#ifndef __REBALANCING_H__
#define __REBALANCING_H__
#include <iostream>
#include "Forwards.hpp"
#include "Term.hpp"
#include "SortHelper.hpp"
#define DEBUG(...)
namespace Kernel {
namespace Rebalancing {
template<class FunctionInverter> class Balancer;
template<class FunctionInverter> class BalanceIter;
struct Node;
class InversionContext;
struct Node {
Node(const Term& term, unsigned index) : _term(&term), _index(index) {}
const Term& term() const { return *_term; }
unsigned index() const { return _index; };
TermList operator*() const {
ASS(inBounds());
return _term->termArg(_index);
}
bool inBounds() const {
return _index < _term->numTermArguments();
}
private:
template<class C>
friend class BalanceIter;
Term const* const _term;
unsigned _index;
};
template<class C>
class Balancer {
const Literal& _lit;
friend class BalanceIter<C>;
public:
Balancer(const Literal& l);
BalanceIter<C> begin() const;
BalanceIter<C> end() const;
};
std::ostream& operator<<(std::ostream& out, const Node&);
template<class C> class BalanceIter {
Stack<Node> _path;
unsigned _litIndex;
const Balancer<C>& _balancer;
friend class Balancer<C>;
BalanceIter(const Balancer<C>&, bool end);
bool inBounds() const;
void findNextVar();
TermList derefPath() const;
void incrementPath();
bool canInvert() const;
public:
void operator++();
const BalanceIter& operator*() const;
template<class D>
friend bool operator!=(const BalanceIter<D>&, const BalanceIter<D>&);
TermList lhs() const;
TermList buildRhs() const;
Literal& build() const;
};
class InversionContext {
const Term& _toInvert;
const TermList _toWrap;
const unsigned _unwrapIdx;
public:
friend std::ostream& operator<<(std::ostream& out, const InversionContext&);
InversionContext(const Term& toInvert, unsigned unwrapIdx, const TermList toWrap) :
_toInvert(toInvert),
_toWrap(toWrap),
_unwrapIdx(unwrapIdx)
{ }
TermList toWrap() const { return _toWrap; }
TermList toUnwrap() const { return _toInvert[_unwrapIdx]; }
const Term& topTerm() const { return _toInvert; }
unsigned topIdx() const { return _unwrapIdx; }
};
template<class C>
Balancer<C>::Balancer(const Literal& l) : _lit(l) { }
template<class C> BalanceIter<C>::BalanceIter(const Balancer<C>& balancer, bool end)
: _path(Stack<Node>())
, _litIndex(end ? 2 : 0)
, _balancer(balancer)
{
if (end) {
DEBUG("end")
} else {
DEBUG("begin(", balancer._lit.toString(), ")");
ASS(balancer._lit.isEquality())
findNextVar();
}
}
template<class C>
BalanceIter<C> Balancer<C>::begin() const {
return BalanceIter<C>(*this, false);
}
template<class C>
BalanceIter<C> Balancer<C>::end() const {
return BalanceIter<C>(*this, true);
}
template<class C> bool BalanceIter<C>::inBounds() const
{
return _litIndex < 2;
}
template<class C> TermList BalanceIter<C>::derefPath() const
{
ASS(_litIndex < 2);
if (_path.isEmpty()) {
return _balancer._lit[_litIndex];
} else {
auto node = _path.top();
return node.term()[node.index()];
}
}
template<class C> bool BalanceIter<C>::canInvert() const
{
if (_path.isEmpty()) {
DEBUG("can invert empty")
return true;
} else {
auto ctxt = InversionContext(_path.top().term(), _path.top().index(), _balancer._lit[1 - _litIndex]);
return C::canInvertTop(ctxt);
}
}
template<class C> void BalanceIter<C>::incrementPath()
{
auto peak = [&]() -> Node& { return _path.top(); };
auto incPeak = [&]() {
++peak()._index;
DEBUG("peakIndex := ", peak().index());
};
auto incLit = [&]() {
_litIndex++;
DEBUG("_litIndex := ", _litIndex)
};
auto inc = [&]() {
if (_path.isEmpty())
incLit();
else
incPeak();
};
do {
if ( derefPath().isTerm()
&& derefPath().term()->numTermArguments() > 0
&& canInvert()) {
DEBUG("push")
_path.push(Node(*derefPath().term(), 0));
} else {
DEBUG("inc")
if (_path.isEmpty()) {
incLit();
} else {
incPeak();
if (!peak().inBounds()) {
do {
_path.pop();
DEBUG("pop()")
inc();
} while (!_path.isEmpty() && !peak().inBounds());
}
}
}
} while (!canInvert() && inBounds());
}
template<class C> void BalanceIter<C>::findNextVar()
{
while(inBounds() && !derefPath().isVar() ) {
incrementPath();
}
}
template<class C> void BalanceIter<C>::operator++() {
incrementPath();
if (inBounds())
findNextVar();
}
template<class C>
const BalanceIter<C>& BalanceIter<C>::operator*() const {
DEBUG(lhs())
return *this;
}
template<class C>
bool operator!=(const BalanceIter<C>& lhs, const BalanceIter<C>& rhs) {
ASS(rhs._path.isEmpty());
auto out = lhs.inBounds(); return out;
}
template<class C>
TermList BalanceIter<C>::lhs() const
{
auto out = derefPath();
ASS_REP(out.isVar(), out);
return out;
}
template<class C>
TermList BalanceIter<C>::buildRhs() const {
ASS(_balancer._lit.numTermArguments() == 2 && _litIndex < 2)
TermList rhs = _balancer._lit[1 - _litIndex];
for (auto n : _path) {
auto ctxt = InversionContext(n.term(), n.index(), rhs);
rhs = C::invertTop(ctxt);
}
return rhs;
}
template<class C>
Literal& BalanceIter<C>::build() const {
return Literal::createEquality(_balancer._lit.polarity(), lhs(), buildRhs(), SortHelper::getEqualityArgumentSort(_balancer._lit));
}
}
} #undef DEBUG
#endif