diff --git a/src/query/plan/cost_estimator.hpp b/src/query/plan/cost_estimator.hpp index 67eef6132..30eceb2d3 100644 --- a/src/query/plan/cost_estimator.hpp +++ b/src/query/plan/cost_estimator.hpp @@ -91,11 +91,11 @@ class CostEstimator : public HierarchicalLogicalOperatorVisitor { using HierarchicalLogicalOperatorVisitor::PostVisit; using HierarchicalLogicalOperatorVisitor::PreVisit; - CostEstimator(TDbAccessor *db_accessor, const Parameters ¶meters) - : db_accessor_(db_accessor), parameters(parameters), scope_(Scope()) {} + CostEstimator(TDbAccessor *db_accessor, const Parameters ¶meters, const SymbolTable &table) + : db_accessor_(db_accessor), parameters(parameters), table_(table), scopes_{Scope()} {} - CostEstimator(TDbAccessor *db_accessor, const Parameters ¶meters, Scope scope) - : db_accessor_(db_accessor), parameters(parameters), scope_(scope) {} + CostEstimator(TDbAccessor *db_accessor, const Parameters ¶meters, const SymbolTable &table, Scope scope) + : db_accessor_(db_accessor), parameters(parameters), table_(table), scopes_{scope} {} bool PostVisit(ScanAll &) override { cardinality_ *= db_accessor_->VerticesCount(); @@ -186,9 +186,13 @@ class CostEstimator : public HierarchicalLogicalOperatorVisitor { // TODO: Cost estimate ScanAllById? bool PostVisit(Expand &expand) override { + const auto &scope = scopes_.back(); auto card_param = CardParam::kExpand; - if (HasStatsFor(expand.input_symbol_)) { - card_param = scope_.symbol_stats[expand.input_symbol_.name()].degree; + + auto stats = GetStatsFor(expand.input_symbol_); + + if (stats.has_value()) { + card_param = stats.value().degree; } cardinality_ *= card_param; @@ -269,6 +273,21 @@ class CostEstimator : public HierarchicalLogicalOperatorVisitor { return false; } + bool PostVisit(Produce &op) override { + auto scope = Scope(); + + for (const auto &symbol : op.ModifiedSymbols(table_)) { + auto stats = GetStatsFor(symbol); + if (stats.has_value()) { + scope.symbol_stats[symbol.name()] = + SymbolStatistics{.cardinality = stats.value().cardinality, .degree = stats.value().degree}; + } + } + + scopes_.push_back(scope); + return true; + } + bool PreVisit(Apply &op) override { double input_cost = EstimateCostOnBranch(&op.input_); double subquery_cost = EstimateCostOnBranch(&op.subquery_); @@ -300,12 +319,13 @@ class CostEstimator : public HierarchicalLogicalOperatorVisitor { // accessor used for cardinality estimates in ScanAll and ScanAllByLabel TDbAccessor *db_accessor_; const Parameters ¶meters; - Scope scope_; + const SymbolTable &table_; + std::vector scopes_; void IncrementCost(double param) { cost_ += param * cardinality_; } double EstimateCostOnBranch(std::shared_ptr *branch) { - CostEstimator cost_estimator(db_accessor_, parameters); + CostEstimator cost_estimator(db_accessor_, parameters, table_); (*branch)->Accept(cost_estimator); return cost_estimator.cost(); } @@ -333,27 +353,39 @@ class CostEstimator : public HierarchicalLogicalOperatorVisitor { return std::nullopt; } - bool HasStatsFor(Symbol &symbol) const { return utils::Contains(scope_.symbol_stats, symbol.name()); } + bool HasStatsFor(const Symbol &symbol) const { return utils::Contains(scopes_.back().symbol_stats, symbol.name()); } - void SaveStatsFor(Symbol &symbol, storage::LabelIndexStats index_stats) { - scope_.symbol_stats[symbol.name()] = SymbolStatistics{ + std::optional GetStatsFor(const Symbol &symbol) { + if (!HasStatsFor(symbol)) { + return std::nullopt; + } + + auto &scope = scopes_.back(); + return scope.symbol_stats[symbol.name()]; + } + + void SaveStatsFor(const Symbol &symbol, storage::LabelIndexStats index_stats) { + scopes_.back().symbol_stats[symbol.name()] = SymbolStatistics{ .cardinality = index_stats.count, .degree = index_stats.avg_degree, }; } - void SaveStatsFor(Symbol &symbol, storage::LabelPropertyIndexStats index_stats) { - scope_.symbol_stats[symbol.name()] = SymbolStatistics{ + void SaveStatsFor(const Symbol &symbol, storage::LabelPropertyIndexStats index_stats) { + scopes_.back().symbol_stats[symbol.name()] = SymbolStatistics{ .cardinality = index_stats.count, .degree = index_stats.avg_degree, }; } + + void DeleteStatsFor(const Symbol &symbol) { scopes_.back().symbol_stats.erase(symbol.name()); } }; /** Returns the estimated cost of the given plan. */ template -double EstimatePlanCost(TDbAccessor *db, const Parameters ¶meters, LogicalOperator &plan) { - CostEstimator estimator(db, parameters); +double EstimatePlanCost(TDbAccessor *db, const Parameters ¶meters, LogicalOperator &plan, + const SymbolTable &table) { + CostEstimator estimator(db, parameters, table); plan.Accept(estimator); return estimator.cost(); } diff --git a/src/query/plan/planner.hpp b/src/query/plan/planner.hpp index cfc89ec73..3f1916d72 100644 --- a/src/query/plan/planner.hpp +++ b/src/query/plan/planner.hpp @@ -47,8 +47,9 @@ class PostProcessor final { } template - double EstimatePlanCost(const std::unique_ptr &plan, TVertexCounts *vertex_counts) { - return query::plan::EstimatePlanCost(vertex_counts, parameters_, *plan); + double EstimatePlanCost(const std::unique_ptr &plan, TVertexCounts *vertex_counts, + SymbolTable &table) { + return query::plan::EstimatePlanCost(vertex_counts, parameters_, *plan, table); } }; @@ -97,7 +98,7 @@ auto MakeLogicalPlan(TPlanningContext *context, TPlanPostProcess *post_process, // Plans are generated lazily and the current plan will disappear, so // it's ok to move it. auto rewritten_plan = post_process->Rewrite(std::move(plan), context); - double cost = post_process->EstimatePlanCost(rewritten_plan, &vertex_counts); + double cost = post_process->EstimatePlanCost(rewritten_plan, &vertex_counts, *context->symbol_table); if (!curr_plan || cost < total_cost) { curr_plan.emplace(std::move(rewritten_plan)); total_cost = cost; @@ -106,7 +107,7 @@ auto MakeLogicalPlan(TPlanningContext *context, TPlanPostProcess *post_process, } else { auto plan = MakeLogicalPlanForSingleQuery(query_parts, context); auto rewritten_plan = post_process->Rewrite(std::move(plan), context); - total_cost = post_process->EstimatePlanCost(rewritten_plan, &vertex_counts); + total_cost = post_process->EstimatePlanCost(rewritten_plan, &vertex_counts, *context->symbol_table); curr_plan.emplace(std::move(rewritten_plan)); } diff --git a/src/query/plan/rewrite/index_lookup.hpp b/src/query/plan/rewrite/index_lookup.hpp index feac431fe..9e4fc7c9a 100644 --- a/src/query/plan/rewrite/index_lookup.hpp +++ b/src/query/plan/rewrite/index_lookup.hpp @@ -505,7 +505,7 @@ class IndexLookupRewriter final : public HierarchicalLogicalOperatorVisitor { // FilterInfo with PropertyFilter. FilterInfo filter; int64_t vertex_count; - std::optional index_stats; + std::optional index_stats; }; bool DefaultPreVisit() override { throw utils::NotYetImplemented("optimizing index lookup"); } @@ -572,8 +572,8 @@ class IndexLookupRewriter final : public HierarchicalLogicalOperatorVisitor { * @param vertex_count: New index's number of vertices. * @return -1 if the new index is better, 0 if they are equal and 1 if the existing one is better. */ - auto compare_indices = [](std::optional &found, std::optional &new_stats, - int vertex_count) { + auto compare_indices = [](std::optional &found, + std::optional &new_stats, int vertex_count) { if (!new_stats.has_value()) { return 0; } @@ -610,7 +610,8 @@ class IndexLookupRewriter final : public HierarchicalLogicalOperatorVisitor { }; int64_t vertex_count = db_->VerticesCount(GetLabel(label), GetProperty(property)); - std::optional new_stats = db_->GetIndexStats(GetLabel(label), GetProperty(property)); + std::optional new_stats = + db_->GetIndexStats(GetLabel(label), GetProperty(property)); // Conditions, from more to less important: // the index with 10x less vertices is better. diff --git a/tests/benchmark/query/planner.cpp b/tests/benchmark/query/planner.cpp index b73258a87..8d898c1b6 100644 --- a/tests/benchmark/query/planner.cpp +++ b/tests/benchmark/query/planner.cpp @@ -131,7 +131,7 @@ static void BM_PlanAndEstimateIndexedMatching(benchmark::State &state) { auto plans = memgraph::query::plan::MakeLogicalPlanForSingleQuery( query_parts, &ctx); for (auto plan : plans) { - memgraph::query::plan::EstimatePlanCost(&dba, parameters, *plan); + memgraph::query::plan::EstimatePlanCost(&dba, parameters, *plan, symbol_table); } } } @@ -161,7 +161,7 @@ static void BM_PlanAndEstimateIndexedMatchingWithCachedCounts(benchmark::State & auto plans = memgraph::query::plan::MakeLogicalPlanForSingleQuery( query_parts, &ctx); for (auto plan : plans) { - memgraph::query::plan::EstimatePlanCost(&vertex_counts, parameters, *plan); + memgraph::query::plan::EstimatePlanCost(&vertex_counts, parameters, *plan, symbol_table); } } } diff --git a/tests/manual/interactive_planning.cpp b/tests/manual/interactive_planning.cpp index 1a1909ba3..a8b66ff26 100644 --- a/tests/manual/interactive_planning.cpp +++ b/tests/manual/interactive_planning.cpp @@ -463,7 +463,7 @@ auto MakeLogicalPlans(memgraph::query::CypherQuery *query, memgraph::query::AstS memgraph::query::AstStorage ast_copy; auto unoptimized_plan = plan->Clone(&ast_copy); auto rewritten_plan = post_process.Rewrite(std::move(plan), &ctx); - double cost = post_process.EstimatePlanCost(rewritten_plan, dba); + double cost = post_process.EstimatePlanCost(rewritten_plan, dba, symbol_table); interactive_plans.push_back( InteractivePlan{std::move(unoptimized_plan), std::move(ast_copy), std::move(rewritten_plan), cost}); } diff --git a/tests/unit/query_cost_estimator.cpp b/tests/unit/query_cost_estimator.cpp index 2253ce0f2..c1b457099 100644 --- a/tests/unit/query_cost_estimator.cpp +++ b/tests/unit/query_cost_estimator.cpp @@ -74,7 +74,7 @@ class QueryCostEstimator : public ::testing::Test { } auto Cost() { - CostEstimator cost_estimator(&*dba, parameters_); + CostEstimator cost_estimator(&*dba, parameters_, symbol_table_); last_op_->Accept(cost_estimator); return cost_estimator.cost(); }