Skip to content
Draft
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
70 changes: 70 additions & 0 deletions jlm/rvsdg/region.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -765,6 +765,76 @@ RegionObserver::RegionObserver(const Region & region)
region.observers_ = this;
}

class RegionTreeObserver::SingleRegionObserver final : public RegionObserver
{
public:
~SingleRegionObserver() noexcept override = default;

SingleRegionObserver(const Region & region, RegionTreeObserver & regionTreeObserver)
: RegionObserver(region),
regionTreeObserver_(&regionTreeObserver)
{}

void
onNodeCreate(Node * node) override
{
regionTreeObserver_->onNodeCreate(node);
}

void
onNodeDestroy(Node * node) override
{
regionTreeObserver_->onNodeDestroy(node);
}

void
onInputCreate(Input * input) override
{
regionTreeObserver_->onInputCreate(input);
}

void
onInputChange(Input * input, Output * old_origin, Output * new_origin) override
{
regionTreeObserver_->onInputChange(input, old_origin, new_origin);
}

void
onInputDestroy(Input * input) override
{
regionTreeObserver_->onInputDestroy(input);
}

private:
RegionTreeObserver * regionTreeObserver_;
};

RegionTreeObserver::~RegionTreeObserver() noexcept = default;

RegionTreeObserver::RegionTreeObserver(const Region & rootRegion)
{
createSingleRegionObservers(rootRegion);
}

void
RegionTreeObserver::createSingleRegionObservers(const Region & region)
{
for (auto & node : region.Nodes())
{
// Handle innermost regions first
if (const auto structuralNode = dynamic_cast<const StructuralNode *>(&node))
{
for (auto & subregion : structuralNode->Subregions())
{
createSingleRegionObservers(subregion);
}
}

auto singleRegionObserver = std::make_unique<SingleRegionObserver>(region, *this);
regionObservers_.emplace_back(std::move(singleRegionObserver));
}
}

std::unordered_map<const Node *, size_t>
computeDepthMap(const Region & region)
{
Expand Down
69 changes: 69 additions & 0 deletions jlm/rvsdg/region.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -1018,6 +1018,75 @@ class RecordingObserver final : public RegionObserver
std::vector<size_t> destroyedInputIndices_{};
};

/**
* \brief Proxy object to observe changes to a region tree.
*
* Subscribers can implement and instantiate this interface for
* a region tree to receive notifications about the regions in the tree.
*
*/
class RegionTreeObserver
{
class SingleRegionObserver;

public:
virtual ~RegionTreeObserver() noexcept;

explicit RegionTreeObserver(const Region & rootRegion);

RegionTreeObserver(const RegionTreeObserver &) = delete;

RegionTreeObserver &
operator=(const RegionTreeObserver &) = delete;

/**
* Called right after a node is added to the region tree,
* after the node has its inputs and output added.
* @param node the node being added
*/
virtual void
onNodeCreate(Node * node) = 0;

/**
* Called right before a node is removed from the region tree,
* before the node has its inputs and outputs removed.
* @param node the node being removed
*/
virtual void
onNodeDestroy(Node * node) = 0;

/**
* Called after a node gets a new input, or a region in the tree gets a new result.
* This method is not called when creating new nodes, only modifying existing nodes.
* @param input the new input
*/
virtual void
onInputCreate(Input * input) = 0;

/**
* Called right after the given input gets a new origin.
* @param input the input.
* @param old_origin the input's old origin.
* @param new_origin the input's new origin.
*/
virtual void
onInputChange(Input * input, Output * old_origin, Output * new_origin) = 0;

/**
* Called right before a node input or region result is removed in the tree.
* This method is not called when deleting nodes, only modifying existing nodes.
* @param input the input that is removed
*/
virtual void
onInputDestroy(Input * input) = 0;

private:
void
createSingleRegionObservers(const Region & region);

std::vector<std::unique_ptr<SingleRegionObserver>> regionObservers_{};
};

/**
* Computes the depth for all nodes in \p region.
*
Expand Down
Loading