#include "optimizer/agg_key_dependency_optimizer.h"
#include "binder/expression/aggregate_function_expression.h"
#include "binder/expression/expression_util.h"
#include "binder/expression/property_expression.h"
#include "function/aggregate_function.h"
#include "planner/operator/logical_aggregate.h"
#include "planner/operator/logical_distinct.h"
using namespace lbug::binder;
using namespace lbug::common;
using namespace lbug::planner;
namespace lbug {
namespace optimizer {
void AggKeyDependencyOptimizer::rewrite(planner::LogicalPlan* plan) {
visitOperator(plan->getLastOperator().get());
}
void AggKeyDependencyOptimizer::visitOperator(planner::LogicalOperator* op) {
for (auto i = 0u; i < op->getNumChildren(); ++i) {
visitOperator(op->getChild(i).get());
}
visitOperatorSwitch(op);
}
void AggKeyDependencyOptimizer::visitAggregate(planner::LogicalOperator* op) {
auto agg = (LogicalAggregate*)op;
auto [keys, dependentKeys] = resolveKeysAndDependentKeys(agg->getKeys());
agg->setKeys(keys);
if (op->getNumChildren() == 1 && op->getChild(0)->getSchema() != nullptr) {
auto& inSchema = *op->getChild(0)->getSchema();
std::unordered_set<std::string> payloadNames;
for (auto& dep : dependentKeys) {
payloadNames.insert(dep->getUniqueName());
}
for (auto& existing : agg->getDependentKeys()) {
if (!payloadNames.contains(existing->getUniqueName()) &&
inSchema.isExpressionInScope(*existing)) {
payloadNames.insert(existing->getUniqueName());
dependentKeys.push_back(existing);
}
}
}
agg->setDependentKeys(dependentKeys);
for (auto& aggExpr : agg->getAggregates()) {
if (aggExpr->expressionType != ExpressionType::AGGREGATE_FUNCTION) {
continue;
}
auto& aggFunc = aggExpr->constCast<AggregateFunctionExpression>();
if (aggFunc.getFunction().name != function::CollectFunction::name ||
aggFunc.getNumChildren() != 1) {
continue;
}
std::unordered_set<std::string> groupNames;
for (auto& key : agg->getKeys()) {
groupNames.insert(key->getUniqueName());
}
collectVarDeps[aggExpr->getUniqueName()] = std::move(groupNames);
}
}
void AggKeyDependencyOptimizer::visitDistinct(planner::LogicalOperator* op) {
auto distinct = (LogicalDistinct*)op;
auto [keys, dependentKeys] = resolveKeysAndDependentKeys(distinct->getKeys());
distinct->setKeys(keys);
distinct->setPayloads(dependentKeys);
}
std::pair<binder::expression_vector, binder::expression_vector>
AggKeyDependencyOptimizer::resolveKeysAndDependentKeys(const expression_vector& inputKeys) {
std::unordered_set<std::string> primaryVarNames;
for (auto& key : inputKeys) {
if (key->expressionType == ExpressionType::PROPERTY) {
auto property = (PropertyExpression*)key.get();
if (property->isPrimaryKey() || property->isInternalID()) {
primaryVarNames.insert(property->getVariableName());
}
}
}
std::unordered_set<std::string> inputKeyNames;
for (auto& key : inputKeys) {
inputKeyNames.insert(key->getUniqueName());
}
binder::expression_vector keys;
binder::expression_vector dependentKeys;
for (auto& key : inputKeys) {
auto collectIt = collectVarDeps.find(key->getUniqueName());
if (collectIt != collectVarDeps.end() && !collectIt->second.empty()) {
bool determined = true;
for (auto& groupName : collectIt->second) {
if (!inputKeyNames.contains(groupName)) {
determined = false;
break;
}
}
if (determined) {
dependentKeys.push_back(key);
continue;
}
}
if (key->expressionType == ExpressionType::PROPERTY) {
auto property = (PropertyExpression*)key.get();
if (property->isPrimaryKey() ||
property->isInternalID()) { keys.push_back(key);
} else if (primaryVarNames.contains(property->getVariableName())) {
dependentKeys.push_back(key);
} else {
keys.push_back(key);
}
} else if (ExpressionUtil::isNodePattern(*key) || ExpressionUtil::isRelPattern(*key)) {
if (primaryVarNames.contains(key->getUniqueName())) {
dependentKeys.push_back(key);
} else {
keys.push_back(key);
}
} else {
keys.push_back(key);
}
}
return std::make_pair(std::move(keys), std::move(dependentKeys));
}
} }