diff --git a/extension/json/test/lsqb_queries_json.test b/extension/json/test/lsqb_queries_json.test index 79a5346fb81..f7998643384 100644 --- a/extension/json/test/lsqb_queries_json.test +++ b/extension/json/test/lsqb_queries_json.test @@ -1,5 +1,5 @@ -DATASET JSON CSV_TO_JSON(lsqb-sf01) --BUFFER_POOL_SIZE 1073741824 +-BUFFER_POOL_SIZE 4294967296 -- -CASE LSQBTestJSON diff --git a/src/include/optimizer/cardinality_updater.h b/src/include/optimizer/cardinality_updater.h index 50d6f99d331..22f44eceed8 100644 --- a/src/include/optimizer/cardinality_updater.h +++ b/src/include/optimizer/cardinality_updater.h @@ -27,6 +27,7 @@ class CardinalityUpdater : public LogicalOperatorVisitor { void visitOperatorDefault(planner::LogicalOperator* op); void visitScanNodeTable(planner::LogicalOperator* op) override; void visitExtend(planner::LogicalOperator* op) override; + void visitRecursiveExtend(planner::LogicalOperator* op) override; void visitHashJoin(planner::LogicalOperator* op) override; void visitCrossProduct(planner::LogicalOperator* op) override; void visitIntersect(planner::LogicalOperator* op) override; diff --git a/src/optimizer/cardinality_updater.cpp b/src/optimizer/cardinality_updater.cpp index 96ec7dd2b0e..6e4275e13d6 100644 --- a/src/optimizer/cardinality_updater.cpp +++ b/src/optimizer/cardinality_updater.cpp @@ -2,6 +2,7 @@ #include "planner/join_order/cardinality_estimator.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_filter.h" #include "planner/operator/logical_flatten.h" @@ -19,6 +20,8 @@ void CardinalityUpdater::visitOperator(planner::LogicalOperator* op) { for (auto i = 0u; i < op->getNumChildren(); ++i) { visitOperator(op->getChild(i).get()); } + // we need to recompute the cardinality multipliers for each factorized group + op->computeFactorizedSchema(); visitOperatorSwitchWithDefault(op); } @@ -32,6 +35,10 @@ void CardinalityUpdater::visitOperatorSwitchWithDefault(planner::LogicalOperator visitExtend(op); break; } + case planner::LogicalOperatorType::RECURSIVE_EXTEND: { + visitRecursiveExtend(op); + break; + } case planner::LogicalOperatorType::HASH_JOIN: { visitHashJoin(op); break; @@ -83,6 +90,18 @@ void CardinalityUpdater::visitExtend(planner::LogicalOperator* op) { const auto extensionRate = cardinalityEstimator.getExtensionRate(*extend.getRel(), *extend.getBoundNode(), transaction); extend.setCardinality(cardinalityEstimator.estimateExtend(extensionRate, *op->getChild(0))); + auto group = extend.getSchema()->getGroup(extend.getNbrNode()->getInternalID()); + group->setMultiplier(extensionRate); +} + +void CardinalityUpdater::visitRecursiveExtend(planner::LogicalOperator* op) { + KU_ASSERT(transaction); + auto& extend = op->cast(); + const auto extensionRate = cardinalityEstimator.getExtensionRate(*extend.getRel(), + *extend.getBoundNode(), transaction); + extend.setCardinality(cardinalityEstimator.estimateExtend(extensionRate, *op->getChild(0))); + auto group = extend.getSchema()->getGroup(extend.getNbrNode()->getInternalID()); + group->setMultiplier(extensionRate); } void CardinalityUpdater::visitHashJoin(planner::LogicalOperator* op) { diff --git a/src/planner/plan/plan_read.cpp b/src/planner/plan/plan_read.cpp index fa640d8fa95..0598350a9d7 100644 --- a/src/planner/plan/plan_read.cpp +++ b/src/planner/plan/plan_read.cpp @@ -146,7 +146,14 @@ void Planner::planGDSCall(const BoundReadingClause& readingClause, gdsCall->computeFactorizedSchema(); probePlan.setLastOperator(gdsCall); if (gdsCall->constPtrCast()->getInfo().func.name == "QFTS") { - auto op = plan->getLastOperator()->getChild(0); + // Hack to deal with join ordering + auto scanParent = plan->getLastOperator(); + KU_ASSERT(scanParent->getNumChildren() >= 2); + idx_t scanChildIdx = (scanParent->getChild(1)->getOperatorType() == + LogicalOperatorType::SCAN_NODE_TABLE) ? + 1 : + 0; + auto op = scanParent->getChild(scanChildIdx); auto prop = bindData->getNodeInput()->constCast().getPropertyExpression( "df"); diff --git a/src/planner/plan/plan_subquery.cpp b/src/planner/plan/plan_subquery.cpp index f4aa77ada29..09e1a0c509f 100644 --- a/src/planner/plan/plan_subquery.cpp +++ b/src/planner/plan/plan_subquery.cpp @@ -1,6 +1,7 @@ #include "binder/expression/expression_util.h" #include "binder/expression/subquery_expression.h" #include "binder/expression_visitor.h" +#include "planner/join_order/cost_model.h" #include "planner/operator/factorization/flatten_resolver.h" #include "planner/planner.h" @@ -164,6 +165,17 @@ void Planner::planOptionalMatch(const QueryGraphCollection& queryGraphCollection } } +template AppendJoinFunc, + std::invocable EstimateJoinCostFunc> +static void planRegularMatchJoinOrder(LogicalPlan& leftPlan, LogicalPlan& rightPlan, + const AppendJoinFunc& appendJoinFunc, const EstimateJoinCostFunc& estimateJoinCostFunc) { + if (estimateJoinCostFunc(leftPlan, rightPlan) <= estimateJoinCostFunc(rightPlan, leftPlan)) { + appendJoinFunc(leftPlan, rightPlan, leftPlan); + } else { + appendJoinFunc(rightPlan, leftPlan, leftPlan); + } +} + void Planner::planRegularMatch(const QueryGraphCollection& queryGraphCollection, const expression_vector& predicates, LogicalPlan& leftPlan) { expression_vector predicatesToPushDown, predicatesToPullUp; @@ -188,7 +200,15 @@ void Planner::planRegularMatch(const QueryGraphCollection& queryGraphCollection, if (leftPlan.hasUpdate()) { appendCrossProduct(*rightPlan, leftPlan, leftPlan); } else { - appendCrossProduct(leftPlan, *rightPlan, leftPlan); + planRegularMatchJoinOrder( + leftPlan, *rightPlan, + [this](LogicalPlan& leftPlan, LogicalPlan& rightPlan, LogicalPlan& resultPlan) { + appendCrossProduct(leftPlan, rightPlan, resultPlan); + }, + [](LogicalPlan&, LogicalPlan& rightPlan) { + // we want to minimize the cardinality of the build plan + return rightPlan.getCardinality(); + }); } } else { // TODO(Xiyang): there is a question regarding if we want to plan as a correlated subquery @@ -201,7 +221,15 @@ void Planner::planRegularMatch(const QueryGraphCollection& queryGraphCollection, if (leftPlan.hasUpdate()) { appendHashJoin(joinNodeIDs, JoinType::INNER, *rightPlan, leftPlan, leftPlan); } else { - appendHashJoin(joinNodeIDs, JoinType::INNER, leftPlan, *rightPlan, leftPlan); + planRegularMatchJoinOrder( + leftPlan, *rightPlan, + [this, &joinNodeIDs](LogicalPlan& leftPlan, LogicalPlan& rightPlan, + LogicalPlan& resultPlan) { + appendHashJoin(joinNodeIDs, JoinType::INNER, leftPlan, rightPlan, resultPlan); + }, + [&joinNodeIDs](LogicalPlan& leftPlan, LogicalPlan& rightPlan) { + return CostModel::computeHashJoinCost(joinNodeIDs, leftPlan, rightPlan); + }); } } for (auto& predicate : predicatesToPullUp) { diff --git a/test/planner/CMakeLists.txt b/test/planner/CMakeLists.txt index 99b658493dc..d746f1f20f7 100644 --- a/test/planner/CMakeLists.txt +++ b/test/planner/CMakeLists.txt @@ -1 +1,3 @@ -add_kuzu_test(planner_tests cardinality_test.cpp) +add_kuzu_test(planner_tests + cardinality_test.cpp + planner_test.cpp) diff --git a/test/planner/cardinality_test.cpp b/test/planner/cardinality_test.cpp index deb836983ac..8af4f3d6947 100644 --- a/test/planner/cardinality_test.cpp +++ b/test/planner/cardinality_test.cpp @@ -102,6 +102,18 @@ TEST_F(CardinalityTest, TestOperators) { EXPECT_GT(joinOp->getCardinality(), 1); } + // Recursive Extend + { + conn->query("CALL enable_gds=false"); + auto plan = getRoot("EXPLAIN LOGICAL MATCH (a)-[r:knows*1..2]-(b:person) WHERE " + "a.ID = 9 AND b.ID = 10 RETURN COUNT(*)"); + auto* extendOp = getOpWithType(plan->getLastOperator().get(), + planner::LogicalOperatorType::RECURSIVE_EXTEND); + ASSERT_NE(nullptr, extendOp); + EXPECT_EQ(200, extendOp->getCardinality()); + conn->query("CALL enable_gds=true"); + } + // Intersect + Flatten { auto plan = getRoot( @@ -112,7 +124,7 @@ TEST_F(CardinalityTest, TestOperators) { auto* intersect = getOpWithType(plan->getLastOperator().get(), planner::LogicalOperatorType::INTERSECT); ASSERT_NE(nullptr, intersect); - EXPECT_EQ(intersect->getCardinality(), 1); + EXPECT_EQ(intersect->getCardinality(), 2); auto* flatten = getOpWithType(plan->getLastOperator().get(), planner::LogicalOperatorType::FLATTEN); diff --git a/test/planner/planner_test.cpp b/test/planner/planner_test.cpp new file mode 100644 index 00000000000..b6025ece9e9 --- /dev/null +++ b/test/planner/planner_test.cpp @@ -0,0 +1,65 @@ +#include "graph_test/graph_test.h" +#include "planner/operator/logical_plan_util.h" +#include "test_runner/test_runner.h" + +namespace kuzu { +namespace testing { + +class PlannerTest : public DBTest { +public: + std::string getInputDir() override { + return TestHelper::appendKuzuRootPath("dataset/tinysnb/"); + } + + std::string getEncodedPlan(const std::string& query) { + return planner::LogicalPlanUtil::encodeJoin(*getRoot(query)); + } + std::unique_ptr getRoot(const std::string& query) { + return TestRunner::getLogicalPlan(query, *conn); + } + std::pair getSource( + planner::LogicalOperator* op, planner::LogicalOperator* parent = nullptr) { + if (op->getNumChildren() == 0) { + return {parent, op}; + } + return getSource(op->getChild(0).get(), op); + } + planner::LogicalOperator* getOpWithType(planner::LogicalOperator* op, + planner::LogicalOperatorType type) { + if (op->getOperatorType() == type) { + return op; + } + if (op->getNumChildren() == 0) { + return nullptr; + } + return getOpWithType(op->getChild(0).get(), type); + } +}; + +TEST_F(PlannerTest, TestSubqueryJoinOrder) { + // Cross Product + { + // We should pick the smaller table to be on the build side + auto query = "MATCH (a:person) WITH a MATCH (b:organisation) RETURN *"; + EXPECT_STREQ("CP(){S(a)}{S(b)}", getEncodedPlan(query).c_str()); + auto queryFlipped = "MATCH (b:organisation) WITH b MATCH (a:person) RETURN *"; + EXPECT_STREQ("CP(){S(a)}{S(b)}", getEncodedPlan(queryFlipped).c_str()); + } + + // Hash Join + { + // cardinality(person) > cardinality(studyAt) + // scan(person) should go on probe side + auto query = "MATCH (a:person) WITH a MATCH (a)-[s:studyAt]->(b:organisation) RETURN *"; + EXPECT_STREQ("HJ(a._ID){S(a)}{HJ(b._ID){S(b)}{E(b)S(a)}}", getEncodedPlan(query).c_str()); + + // cardinality(organisation) < cardinality(studyAt) + cardinality(workAt) + // scan(organisation) should go on build side + auto queryFlipped = + "MATCH (b:organisation) WITH b MATCH (a:person)-[s:studyAt|:workAt]->(b) RETURN *"; + EXPECT_STREQ("HJ(b._ID){E(b)S(a)}{S(b)}", getEncodedPlan(queryFlipped).c_str()); + } +} + +} // namespace testing +} // namespace kuzu diff --git a/test/test_files/agg/multi_query_part.test b/test/test_files/agg/multi_query_part.test index e89336b5fdb..cc812bc6c96 100644 --- a/test/test_files/agg/multi_query_part.test +++ b/test/test_files/agg/multi_query_part.test @@ -96,3 +96,13 @@ Bob|Dan|3 Dan|Alice|3 Dan|Bob|3 Dan|Carol|3 + +# If the aggregation happens on the build side of the join +# A flatten will needed to be appended to the probe side +# Otherwise we will have the incorrect number of results +-LOG GroupByMultiQueryTest8 +-STATEMENT MATCH (a:person)-[k1:knows]->(b:person) WITH a, b, COUNT(*) AS s + MATCH (c:person)-[k2:knows]->(a) WITH c, s + MATCH (c) WHERE c.ID = 0 RETURN COUNT(*) +---- 1 +9 diff --git a/test/test_files/lsqb/lsqb_queries.test b/test/test_files/lsqb/lsqb_queries.test index 84dca9a0845..3654ca181cb 100644 --- a/test/test_files/lsqb/lsqb_queries.test +++ b/test/test_files/lsqb/lsqb_queries.test @@ -1,5 +1,5 @@ -DATASET CSV lsqb-sf01 --BUFFER_POOL_SIZE 1073741824 +-BUFFER_POOL_SIZE 4294967296 -- diff --git a/test/test_files/lsqb/lsqb_queries_parquet.test b/test/test_files/lsqb/lsqb_queries_parquet.test index f4c5f60a092..fbac68dafe7 100644 --- a/test/test_files/lsqb/lsqb_queries_parquet.test +++ b/test/test_files/lsqb/lsqb_queries_parquet.test @@ -1,5 +1,5 @@ -DATASET PARQUET CSV_TO_PARQUET(lsqb-sf01) --BUFFER_POOL_SIZE 1073741824 +-BUFFER_POOL_SIZE 4294967296 -- -CASE LSQBTestParquet