Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions src/include/optimizer/count_rel_table_optimizer.h
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,13 @@ class CountRelTableOptimizer : public LogicalOperatorVisitor {
std::shared_ptr<planner::LogicalOperator> tryRewriteExtendChainCount(
std::shared_ptr<planner::LogicalOperator> op);

// Rewrite COUNT(*) over the LSQB-q9 shape: a path n0-R-n1-R-n2-... whose first two hops
// are BOTH-direction R extends with an anti-edge filter NOT(EXISTS{n0-R-n2}) (+ optional
// id(n0)<>id(n2)), plus a filter-free suffix chain from n2. Replaced with count arithmetic
// T - A - S (see CountAntiEdgeChain).
std::shared_ptr<planner::LogicalOperator> tryRewriteAntiEdgeChainCount(
std::shared_ptr<planner::LogicalOperator> op);

// Check if the aggregate is a simple (non-distinct) COUNT with no keys.
bool isSimpleCount(planner::LogicalOperator* op) const;

Expand Down
1 change: 1 addition & 0 deletions src/include/planner/operator/logical_operator.h
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ enum class LogicalOperatorType : uint8_t {
ATTACH_DATABASE,
COPY_FROM,
COPY_TO,
COUNT_ANTI_EDGE_CHAIN,
COUNT_EXTEND_CHAIN,
COUNT_REL_TABLE,
CREATE_GRAPH,
Expand Down
116 changes: 116 additions & 0 deletions src/include/planner/operator/scan/logical_count_anti_edge_chain.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,116 @@
#pragma once

#include <string>
#include <vector>

#include "binder/expression/expression.h"
#include "catalog/catalog_entry/rel_group_catalog_entry.h"
#include "common/enums/extend_direction.h"
#include "common/enums/rel_direction.h"
#include "common/types/types.h"
#include "planner/operator/logical_operator.h"
#include "planner/operator/scan/logical_count_extend_chain.h"

namespace lbug {
namespace planner {

struct LogicalCountAntiEdgeChainPrintInfo final : OPPrintInfo {
std::string relTableName;
uint64_t numSuffixHops;

LogicalCountAntiEdgeChainPrintInfo(std::string relTableName, uint64_t numSuffixHops)
: relTableName{std::move(relTableName)}, numSuffixHops{numSuffixHops} {}

std::string toString() const override {
return "Anti-edge rel: " + relTableName + ", suffix hops: " + std::to_string(numSuffixHops);
}

std::unique_ptr<OPPrintInfo> copy() const override {
return std::make_unique<LogicalCountAntiEdgeChainPrintInfo>(relTableName, numSuffixHops);
}
};

/**
* LogicalCountAntiEdgeChain computes COUNT(*) over a path whose first two hops h1, h2 (same rel
* table R, same node table) are additionally filtered by an anti-edge predicate:
* NOT(EXISTS{MATCH (n0)-[:R]-(n2)}) and optionally id(n0) <> id(n2). Detected by
* CountRelTableOptimizer::tryRewriteAntiEdgeChainCount from plans shaped like LSQB q9:
*
* HASH_JOIN[INNER key=n2]
* FILTER[NOT(EXISTS{...})]
* HASH_JOIN[MARK keys={n0,n1...}] (anti-edge rows as one side)
* chain of R-extends from scan(n1) + anti-edge R-extend over scan(n0)
* suffix chain of extends from scan(n2)
*
* The two chain hops enumerate rows in directions chainN0Dir/chainN2Dir (FWD, BWD or BOTH) and
* the anti-edge match enumerates rows in antiEdgeDir. The count is computed with pure count
* arithmetic (see CountAntiEdgeChain):
*
* count = T - A - S
* T = sum over chain rows (n1,n2) of degChainN0(n1)*D(n2) (all tuples)
* A = sum over anti-edge rows (n0,n2) of N'(n0,n2)*D(n2) (anti-edge present)
* S = sum over mid nodes n1, nodes p of multN0(p)*multN2(p)*D(p) (n0=n2 tuples; only
* when the id<> filter
* is present)
* where D = per-node suffix path counts, degChainN0(n1) = number of n0-hop rows at n1,
* multN0/multN2(p) = number of n0/n2-hop rows between n1 and p, and
* N'(n0,n2) = sum over mid nodes x of multRevN0(n0,x)*multRevN2(n2,x) = number of chain
* enumerations connecting n0 to n2.
*/
class LogicalCountAntiEdgeChain final : public LogicalOperator {
static constexpr LogicalOperatorType type_ = LogicalOperatorType::COUNT_ANTI_EDGE_CHAIN;

public:
LogicalCountAntiEdgeChain(std::vector<CountChainHop> suffixHops,
catalog::RelGroupCatalogEntry* antiRelEntry,
std::vector<common::table_id_t> antiRelTableIDs, common::table_id_t midNodeTableID,
common::ExtendDirection chainN0Dir, common::ExtendDirection chainN2Dir,
common::ExtendDirection antiEdgeDir, bool hasNotEquals,
std::shared_ptr<binder::Expression> countExpr)
: LogicalOperator{type_}, suffixHops{std::move(suffixHops)}, antiRelEntry{antiRelEntry},
antiRelTableIDs{std::move(antiRelTableIDs)}, midNodeTableID{midNodeTableID},
chainN0Dir{chainN0Dir}, chainN2Dir{chainN2Dir}, antiEdgeDir{antiEdgeDir},
hasNotEquals{hasNotEquals}, countExpr{std::move(countExpr)} {
cardinality = 1;
}

void computeFactorizedSchema() override;
void computeFlatSchema() override;

std::string getExpressionsForPrinting() const override { return countExpr->toString(); }

const std::vector<CountChainHop>& getSuffixHops() const { return suffixHops; }
catalog::RelGroupCatalogEntry* getAntiRelEntry() const { return antiRelEntry; }
const std::vector<common::table_id_t>& getAntiRelTableIDs() const { return antiRelTableIDs; }
common::table_id_t getMidNodeTableID() const { return midNodeTableID; }
common::ExtendDirection getChainN0Dir() const { return chainN0Dir; }
common::ExtendDirection getChainN2Dir() const { return chainN2Dir; }
common::ExtendDirection getAntiEdgeDir() const { return antiEdgeDir; }
bool getHasNotEquals() const { return hasNotEquals; }
std::shared_ptr<binder::Expression> getCountExpr() const { return countExpr; }

std::unique_ptr<OPPrintInfo> getPrintInfo() const override {
return std::make_unique<LogicalCountAntiEdgeChainPrintInfo>(antiRelEntry->getName(),
suffixHops.size());
}

std::unique_ptr<LogicalOperator> copy() override {
return std::make_unique<LogicalCountAntiEdgeChain>(suffixHops, antiRelEntry,
antiRelTableIDs, midNodeTableID, chainN0Dir, chainN2Dir, antiEdgeDir, hasNotEquals,
countExpr);
}

private:
std::vector<CountChainHop> suffixHops;
catalog::RelGroupCatalogEntry* antiRelEntry;
std::vector<common::table_id_t> antiRelTableIDs;
common::table_id_t midNodeTableID;
common::ExtendDirection chainN0Dir;
common::ExtendDirection chainN2Dir;
common::ExtendDirection antiEdgeDir;
bool hasNotEquals;
std::shared_ptr<binder::Expression> countExpr;
};

} // namespace planner
} // namespace lbug
1 change: 1 addition & 0 deletions src/include/processor/operator/physical_operator.h
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@ enum class PhysicalOperatorType : uint8_t {
ATTACH_DATABASE,
BATCH_INSERT,
COPY_TO,
COUNT_ANTI_EDGE_CHAIN,
COUNT_EXTEND_CHAIN,
COUNT_REL_TABLE,
CREATE_GRAPH,
Expand Down
114 changes: 114 additions & 0 deletions src/include/processor/operator/scan/count_anti_edge_chain.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,114 @@
#pragma once

#include <unordered_map>

#include "common/enums/extend_direction.h"
#include "common/enums/rel_direction.h"
#include "common/system_config.h"
#include "common/types/types.h"
#include "processor/data_pos.h"
#include "processor/operator/physical_operator.h"
#include "storage/table/node_table.h"
#include "storage/table/rel_table.h"

namespace lbug {
namespace processor {

struct CountAntiEdgeChainPrintInfo final : OPPrintInfo {
std::string relTableName;
uint64_t numSuffixHops;

CountAntiEdgeChainPrintInfo(std::string relTableName, uint64_t numSuffixHops)
: relTableName{std::move(relTableName)}, numSuffixHops{numSuffixHops} {}

std::string toString() const override {
return "Anti-edge rel: " + relTableName + ", suffix hops: " + std::to_string(numSuffixHops);
}

std::unique_ptr<OPPrintInfo> copy() const override {
return std::make_unique<CountAntiEdgeChainPrintInfo>(relTableName, numSuffixHops);
}
};

/**
* CountAntiEdgeChain computes COUNT(*) over LSQB-q9-shaped queries with pure count arithmetic.
* See LogicalCountAntiEdgeChain for the pattern and the T - A - S formula. No hash joins are
* performed and no tuples are materialized: the operator scans the anti-edge rel table to build
* directional adjacency lists, scans the suffix rel tables backward to propagate per-node
* suffix path counts D, then evaluates T, A and S over flat arrays.
*
* The two chain hops enumerate rows in chainN0Dir/chainN2Dir and the anti-edge match in
* antiEdgeDir (each FWD, BWD or BOTH); the arithmetic models exactly that enumeration.
*
* Single-threaded source emitting one row with the final count (int64).
*/
class CountAntiEdgeChain final : public PhysicalOperator {
static constexpr PhysicalOperatorType type_ = PhysicalOperatorType::COUNT_ANTI_EDGE_CHAIN;

public:
struct Hop {
std::vector<storage::RelTable*> relTables;
std::vector<common::RelDataDirection> scanDirections;
std::vector<storage::NodeTable*> fromNodeTables;
std::vector<storage::NodeTable*> toNodeTables;
};

CountAntiEdgeChain(std::vector<Hop> suffixHops, storage::RelTable* antiRelTable,
storage::NodeTable* midNodeTable, common::ExtendDirection chainN0Dir,
common::ExtendDirection chainN2Dir, common::ExtendDirection antiEdgeDir, bool hasNotEquals,
DataPos countOutputPos, physical_op_id id, std::unique_ptr<OPPrintInfo> printInfo)
: PhysicalOperator{type_, id, std::move(printInfo)}, suffixHops{std::move(suffixHops)},
antiRelTable{antiRelTable}, midNodeTable{midNodeTable}, chainN0Dir{chainN0Dir},
chainN2Dir{chainN2Dir}, antiEdgeDir{antiEdgeDir}, hasNotEquals{hasNotEquals},
countOutputPos{countOutputPos} {}

bool isSource() const override { return true; }
bool isParallel() const override { return false; }

void initLocalStateInternal(ResultSet* resultSet, ExecutionContext* context) override;

bool getNextTuplesInternal(ExecutionContext* context) override;

std::unique_ptr<PhysicalOperator> copy() override {
return std::make_unique<CountAntiEdgeChain>(suffixHops, antiRelTable, midNodeTable,
chainN0Dir, chainN2Dir, antiEdgeDir, hasNotEquals, countOutputPos, id,
printInfo->copy());
}

private:
common::offset_t getOffsetUpperBound(storage::NodeTable* nodeTable) const {
const auto numGroups = nodeTable->getNumNodeGroups();
if (numGroups == 0) {
return 0;
}
return (numGroups - 1) * common::StorageConfig::NODE_GROUP_SIZE +
nodeTable->getNumTuplesInNodeGroup(numGroups - 1);
}

// Scan one rel table in the given direction with sequential bound-node batches, invoking
// the per-batch callback with (scanState, nodeIDVector, nbrVector).
template<typename Func>
void scanRelRows(storage::RelTable* relTable, common::RelDataDirection direction,
storage::NodeTable* boundNodeTable, transaction::Transaction* transaction,
storage::MemoryManager* memoryManager, Func&& callback);

// Backward suffix propagation: returns B[0], the per-node suffix path counts over the mid
// node table.
std::vector<int64_t> computeSuffixCounts(transaction::Transaction* transaction,
storage::MemoryManager* memoryManager);

private:
std::vector<Hop> suffixHops;
storage::RelTable* antiRelTable;
storage::NodeTable* midNodeTable;
common::ExtendDirection chainN0Dir;
common::ExtendDirection chainN2Dir;
common::ExtendDirection antiEdgeDir;
bool hasNotEquals;
DataPos countOutputPos;
common::ValueVector* countVector = nullptr;
bool hasExecuted = false;
};

} // namespace processor
} // namespace lbug
2 changes: 2 additions & 0 deletions src/include/processor/plan_mapper.h
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,8 @@ class PlanMapper {
std::unique_ptr<PhysicalOperator> mapCopyTo(const planner::LogicalOperator* logicalOperator);
std::unique_ptr<PhysicalOperator> mapCountRelTable(
const planner::LogicalOperator* logicalOperator);
std::unique_ptr<PhysicalOperator> mapCountAntiEdgeChain(
const planner::LogicalOperator* logicalOperator);
std::unique_ptr<PhysicalOperator> mapCountExtendChain(
const planner::LogicalOperator* logicalOperator);
std::unique_ptr<PhysicalOperator> mapCreateMacro(
Expand Down
Loading
Loading