#include "Debug/Assertion.hpp"
#include "Lib/Allocator.hpp"
#include "Lib/Environment.hpp"
#include "Lib/ScopedLet.hpp"
#include "Kernel/Clause.hpp"
#include "Kernel/Formula.hpp"
#include "Kernel/Inference.hpp"
#include "Kernel/FormulaUnit.hpp"
#include "Kernel/Problem.hpp"
#include "Kernel/Signature.hpp"
#include "Kernel/Term.hpp"
#include "Kernel/TermIterators.hpp"
#include "Shell/Options.hpp"
#include "Statistics.hpp"
#include "FunctionDefinition.hpp"
#include <algorithm>
#if VDEBUG
#include <iostream>
#endif
namespace Shell {
using namespace std;
using namespace Lib;
using namespace Kernel;
struct FunctionDefinition::Def
{
enum Mark {
UNTOUCHED,
SAFE,
LOOP,
BLOCKED,
UNFOLDED,
REMOVED
};
Clause* defCl;
int fun;
Term* lhs;
Term* rhs;
Mark mark;
bool linear;
bool strict;
bool twoConstDef;
int containedFn;
int examinedArg;
IntList* dependentFns;
bool lhsIsBool(){
return (env.signature->isFoolConstantSymbol(true , fun) ||
env.signature->isFoolConstantSymbol(false, fun));
}
bool* argOccurs;
Def(Term* l, Term* r, bool lin, bool str)
: fun(l->functor()),
lhs(l),
rhs(r),
mark(UNTOUCHED),
linear(lin),
strict(str),
twoConstDef(0),
containedFn(-1),
dependentFns(0),
argOccurs(0)
{
}
~Def()
{
if(argOccurs) {
DEALLOC_KNOWN(argOccurs, lhs->arity()*sizeof(bool), "FunctionDefinition::Def::argOccurs");
}
}
};
void FunctionDefinition::removeUnusedDefinitions(Problem& prb, bool inHigherOrder)
{
if(removeUnusedDefinitions(prb.units(), &prb, inHigherOrder)) {
prb.invalidateByRemoval();
}
}
bool FunctionDefinition::removeUnusedDefinitions(UnitList*& units, Problem* prb, bool inHigherOrder)
{
unsigned funs=env.signature->functions();
Stack<Def*> defStack;
DArray<Def*> def;
DArray<unsigned> occCounter;
def.init(funs, 0);
occCounter.init(funs, 0);
UnitList::DelIterator scanIterator(units);
while(scanIterator.hasNext()) {
Clause* cl=static_cast<Clause*>(scanIterator.next());
unsigned clen=cl->length();
ASS(cl->isClause());
Def* d=isFunctionDefinition(cl,inHigherOrder);
if(d) {
d->defCl=cl;
if(!def[d->fun]) {
defStack.push(d);
def[d->fun]=d;
scanIterator.del();
} else {
delete d;
}
}
for(unsigned i=0;i<clen;i++) {
NonVariableNonTypeIterator nvit((*cl)[i]);
while(nvit.hasNext()) {
unsigned fn=nvit.next()->functor();
occCounter[fn]++;
}
}
}
Stack<Def*> toDo;
Stack<Def*>::Iterator dit(defStack);
while(dit.hasNext()) {
Def* d=dit.next();
unsigned fn=d->fun;
ASS_GE(occCounter[fn],1);
if(occCounter[fn]==1) {
toDo.push(d);
}
}
while(toDo.isNonEmpty()) {
Def* d=toDo.pop();
d->mark=Def::REMOVED;
ASS_EQ(d->defCl->length(), 1);
ASS_EQ(occCounter[d->fun], 1);
NonVariableNonTypeIterator nvit((*d->defCl)[0]);
while(nvit.hasNext()) {
unsigned fn=nvit.next()->functor();
occCounter[fn]--;
if(occCounter[fn]==1 && def[fn]) {
toDo.push(def[fn]);
}
}
ASS_EQ(occCounter[d->fun], 0);
}
bool modified = false;
while(defStack.isNonEmpty()) {
Def* d=defStack.pop();
if(d->mark==Def::REMOVED) {
modified = true;
if(prb) {
ASS_EQ(d->defCl->length(), 1);
prb->addEliminatedFunction(d->fun, (*d->defCl)[0]);
}
}
else {
ASS_EQ(d->mark, Def::UNTOUCHED);
UnitList::push(d->defCl, units);
}
delete d;
}
return modified;
}
void FunctionDefinition::removeAllDefinitions(Problem& prb, bool inHigherOrder)
{
ScopedLet<Problem*> prbLet(_processedProblem, &prb);
if(removeAllDefinitions(prb.units(),inHigherOrder)) {
prb.invalidateByRemoval();
}
}
void FunctionDefinition::reverse(Def* def){
ASS(def->twoConstDef);
Term* temp = def->lhs;
IGNORE_MAYBE_UNINITIALIZED(
def->lhs = def->rhs;
)
def->rhs = temp;
def->fun = def->lhs->functor();
}
bool FunctionDefinition::removeAllDefinitions(UnitList*& units, bool inHigherOrder)
{
UnitList::DelIterator scanIterator(units);
while(scanIterator.hasNext()) {
Clause* cl=static_cast<Clause*>(scanIterator.next());
ASS(cl->isClause());
Def* d=isFunctionDefinition(cl,inHigherOrder);
if(d) {
d->defCl=cl;
bool inserted = false;
if(_defs.insert(d->fun, d)) {
inserted = true;
scanIterator.del();
} else if(_defs.get(d->fun)->twoConstDef){
Def* d2;
ALWAYS(_defs.pop(d->fun, d2));
reverse(d2);
if(!d2->lhsIsBool() && _defs.insert(d2->fun, d2)){
_defs.insert(d->fun, d);
inserted = true;
scanIterator.del();
} else {
reverse(d2); ALWAYS(_defs.insert(d2->fun, d2));
}
} else if(d->twoConstDef){
reverse(d);
if(!d->lhsIsBool() && _defs.insert(d->fun, d)) {
inserted = true;
scanIterator.del();
}
}
if(!inserted){
delete d;
}
}
}
if(!_defs.size()) {
return false;
}
Fn2DefMap::Iterator dit(_defs);
while(dit.hasNext()) {
Def* d=dit.next();
if(d->mark==Def::SAFE || d->mark==Def::BLOCKED) {
continue;
}
ASS(d->mark==Def::UNTOUCHED);
checkDefinitions(d);
}
while(_blockedDefs.isNonEmpty()) {
Def* d=_blockedDefs.pop();
ASS_EQ(d->mark, Def::BLOCKED);
UnitList::push(d->defCl, units);
_defs.remove(d->fun);
delete d;
}
if(_defs.isEmpty()) {
return false;
}
ASS_EQ(_defs.size(), _safeDefs.size());
for(unsigned i=0;i<_safeDefs.size(); i++) {
Def* d=_safeDefs[i];
ASS_EQ(d->mark, Def::SAFE);
d->mark=Def::BLOCKED;
d->defCl=applyDefinitions(d->defCl);
Literal* defEq=(*d->defCl)[0];
if( defEq->nthArgument(0)->term()==d->lhs ) {
d->rhs=defEq->nthArgument(1)->term();
} else {
ASS_EQ(defEq->nthArgument(1)->term(),d->lhs);
d->rhs=defEq->nthArgument(0)->term();
}
d->mark=Def::UNFOLDED;
if(_processedProblem) {
ASS_EQ(d->defCl->length(), 1);
_processedProblem->addEliminatedFunction(d->fun, (*d->defCl)[0]);
}
if (env.options->showPreprocessing()) {
std::cout << "[PP] fn def discovered: "<<(*d->defCl)<<"\n unfolded: "<<(*d->rhs) << std::endl;
}
env.statistics->eliminatedFunctionDefinitions++;
}
UnitList::DelIterator unfoldIterator(units);
while(unfoldIterator.hasNext()) {
Clause* cl=static_cast<Clause*>(unfoldIterator.next());
ASS(cl->isClause());
Clause* newCl=applyDefinitions(cl);
if(cl!=newCl) {
unfoldIterator.replace(newCl);
}
}
_safeDefs.reset();
return true;
}
void FunctionDefinition::checkDefinitions(Def* def0)
{
TermList t=TermList(def0->lhs);
static Stack<TermList*> stack(4);
static Stack<Def*> defCheckingStack(4);
static Stack<Def*> defArgStack(4);
static Stack<Term*> termArgStack(4);
for(;;) {
Def* d;
if(t.isEmpty()) {
d=defCheckingStack.pop();
defArgStack.pop();
termArgStack.pop();
ASS(!d || d->mark==Def::LOOP);
if(d && d->mark==Def::LOOP) {
assignArgOccursData(d);
_safeDefs.push(d);
d->mark=Def::SAFE;
}
} else if(t.isTerm() && !t.term()->isSort()) {
Term* trm=t.term();
Def* checkedDef=0;
toplevel_def:
if(!_defs.find(trm->functor(), d) || d->mark==Def::BLOCKED) {
d=0;
}
if(trm->numTermArguments() > 0 || checkedDef) {
stack.push(trm->termArgs());
defCheckingStack.push(checkedDef);
defArgStack.push(d);
termArgStack.push(trm);
}
if(d) {
if(d->mark==Def::UNTOUCHED) {
d->mark=Def::LOOP;
trm=d->rhs;
checkedDef=d;
goto toplevel_def;
} else if(d->mark==Def::LOOP) {
do{
stack.pop();
defArgStack.pop();
termArgStack.pop();
d=defCheckingStack.pop();
} while(!d);
ASS_EQ(d->mark, Def::LOOP);
d->mark=Def::BLOCKED;
defArgStack.setTop(0);
_blockedDefs.push(d);
} else {
ASS_EQ(d->mark, Def::SAFE);
}
}
}
if(stack.isEmpty()) {
break;
}
TermList* ts=stack.pop();
if(ts->isNonEmpty()) {
Def* argDef=defArgStack.top();
if(argDef) {
ASS_EQ(argDef->mark,Def::SAFE);
Term* parentTerm=termArgStack.top();
while(ts->isNonEmpty() && !argDef->argOccurs[parentTerm->getArgumentIndex(ts)]) {
ts=ts->next();
}
if(ts->isNonEmpty()) {
stack.push(ts->next());
}
} else {
stack.push(ts->next());
}
}
t=*ts;
}
ASS(defCheckingStack.isEmpty());
ASS(defArgStack.isEmpty());
}
void FunctionDefinition::assignArgOccursData(Def* updDef)
{
ASS(!updDef->argOccurs);
if(!updDef->lhs->arity()) {
return;
}
updDef->argOccurs=reinterpret_cast<bool*>(ALLOC_KNOWN(updDef->lhs->arity()*sizeof(bool),
"FunctionDefinition::Def::argOccurs"));
std::memset(updDef->argOccurs, 0, updDef->lhs->arity() * sizeof(bool));
static DHMap<unsigned, unsigned, IdentityHash, DefaultHash> var2argIndex;
var2argIndex.reset();
int argIndex=0;
for (TermList* ts = updDef->lhs->args(); ts->isNonEmpty(); ts=ts->next()) {
int w = ts->var();
var2argIndex.insert(w, argIndex);
argIndex++;
}
TermList t=TermList(updDef->rhs);
static Stack<TermList*> stack(4);
static Stack<Def*> defArgStack(4);
static Stack<Term*> termArgStack(4);
for(;;) {
Def* d;
if(t.isEmpty()) {
defArgStack.pop();
termArgStack.pop();
} else if(t.isTerm()) {
Term* trm=t.term();
if(trm->arity()) {
if(trm->isSort() || !_defs.find(trm->functor(), d) || d->mark==Def::BLOCKED) {
d=0;
}
ASS(!d || d->mark==Def::SAFE);
stack.push(trm->args());
defArgStack.push(d);
termArgStack.push(trm);
}
} else {
ASS(t.isOrdinaryVar());
updDef->argOccurs[var2argIndex.get(t.var())]=true;
}
if(stack.isEmpty()) {
break;
}
TermList* ts=stack.pop();
if(!ts->isEmpty()) {
Def* argDef=defArgStack.top();
if(argDef) {
Term* parentTerm=termArgStack.top();
while(ts->isNonEmpty() && !argDef->argOccurs[parentTerm->getArgumentIndex(ts)]) {
ts=ts->next();
}
if(ts->isNonEmpty()) {
stack.push(ts->next());
}
} else {
stack.push(ts->next());
}
}
t=*ts;
}
}
typedef pair<unsigned,unsigned> BindingSpec;
typedef DHMap<BindingSpec, TermList> BindingMap;
typedef DHMap<BindingSpec, bool> UnfoldedSet;
Term* FunctionDefinition::applyDefinitions(Literal* lit, Stack<Def*>* usedDefs)
{
if (env.options->showPreprocessing()) {
std::cout << "[PP] applying function definitions to literal "<<(*lit) << std::endl;
}
BindingMap bindings;
UnfoldedSet unfolded;
unsigned nextDefIndex=1;
Stack<BindingSpec> bindingsBeingUnfolded;
Stack<unsigned> defIndexes;
Stack<TermList*> toDo;
Stack<Term*> terms;
Stack<bool> modified;
Stack<TermList> args;
bindings.reset();
unfolded.reset();
defIndexes.reset();
toDo.reset();
terms.reset();
modified.reset();
args.reset();
defIndexes.push(0);
modified.push(false);
toDo.push(lit->args());
for(;;) {
TermList* tt=toDo.pop();
if(!tt) {
BindingSpec spec=bindingsBeingUnfolded.pop();
bindings.set(spec, args.top());
unfolded.insert(spec, true);
continue;
}
if(tt->isEmpty()) {
if(terms.isEmpty()) {
ASS(toDo.isEmpty());
break;
}
defIndexes.pop();
Term* orig=terms.pop();
if(!modified.pop()) {
args.truncate(args.length() - orig->arity());
args.push(TermList(orig));
continue;
}
TermList* argLst=&args.top() - (orig->arity()-1);
Term* newTrm;
if(orig->isSort()){
newTrm=AtomicSort::create(static_cast<AtomicSort*>(orig), argLst);
} else {
newTrm=Term::create(orig,argLst);
}
args.truncate(args.length() - orig->arity());
args.push(TermList(newTrm));
modified.setTop(true);
continue;
}
toDo.push(tt->next());
TermList tl=*tt;
unsigned defIndex=defIndexes.top();
Term* t;
if(tl.isVar()) {
ASS(tl.isOrdinaryVar());
if(defIndexes.top()) {
modified.setTop(true);
BindingSpec spec=make_pair(defIndexes.top(), tl.var());
TermList bound=bindings.get(spec);
if(bound.isVar() || unfolded.find(spec)) {
args.push(bound);
continue;
} else {
bindingsBeingUnfolded.push(spec);
toDo.push(0);
defIndex=0;
t=bound.term();
}
} else {
args.push(tl);
continue;
}
} else {
ASS(tl.isTerm());
t=tl.term();
}
Def* d;
if(!t->isSort() && !defIndex && _defs.find(t->functor(), d) && d->mark!=Def::BLOCKED) {
ASS_EQ(d->mark, Def::UNFOLDED);
usedDefs->push(d);
if (env.options->showPreprocessing()) {
std::cout << "[PP] definition of "<<(*t)<<"\n expanded to "<<(*d->rhs) << std::endl;
}
defIndex=nextDefIndex++;
TermList* dargs=d->lhs->args();
TermList* targs=t->args();
while(dargs->isNonEmpty()) {
ASS(targs->isNonEmpty());
ALWAYS(bindings.insert(make_pair(defIndex, dargs->var()), *targs));
dargs=dargs->next();
targs=targs->next();
}
t=d->rhs;
modified.setTop(true);
}
defIndexes.push(defIndex);
terms.push(t);
modified.push(false);
toDo.push(t->args());
}
ASS(toDo.isEmpty());
ASS(terms.isEmpty());
ASS_EQ(modified.length(),1);
ASS_EQ(args.length(),lit->arity());
if(!modified.pop()) {
return lit;
}
TermList* argLst=&args.top() - (lit->arity()-1);
return Literal::create(static_cast<Literal*>(lit),argLst);
}
Clause* FunctionDefinition::applyDefinitions(Clause* cl)
{
unsigned clen=cl->length();
static Stack<Def*> usedDefs(8);
RStack<Literal*> resLits;
ASS(usedDefs.isEmpty());
bool modified=false;
for(unsigned i=0;i<clen;i++) {
Literal* lit=(*cl)[i];
Literal* rlit=static_cast<Literal*>(applyDefinitions(lit, &usedDefs));
resLits->push(rlit);
modified|= rlit!=lit;
}
if(!modified) {
ASS(usedDefs.isEmpty());
return cl;
}
UnitList* premises=0;
std::vector<Term *> extra;
while(usedDefs.isNonEmpty()) {
Def *def = usedDefs.pop();
Clause* defCl=def->defCl;
UnitList::push(defCl, premises);
if(env.options->proofExtra() == Options::ProofExtra::FULL)
extra.push_back(def->lhs);
}
std::reverse(extra.begin(), extra.end());
UnitList::push(cl, premises);
auto res = Clause::fromStack(*resLits, NonspecificInferenceMany(InferenceRule::DEFINITION_UNFOLDING, premises));
if(env.options->proofExtra() == Options::ProofExtra::FULL)
env.proofExtra.insert(res, new FunctionDefinitionExtra(std::move(extra)));
res->setAge(cl->age()); return res;
}
FunctionDefinition::~FunctionDefinition ()
{
Fn2DefMap::Iterator dit(_defs);
while(dit.hasNext()) {
delete dit.next();
}
}
FunctionDefinition::Def*
FunctionDefinition::isFunctionDefinition (Unit& unit, bool inHigherOrder)
{
if(unit.derivedFromGoal() && env.options->ignoreConjectureInPreprocessing()){
return 0;
}
if (unit.isClause()) {
return isFunctionDefinition(static_cast<Clause*>(&unit), inHigherOrder);
}
return isFunctionDefinition(static_cast<FormulaUnit&>(unit), inHigherOrder);
}
FunctionDefinition::Def*
FunctionDefinition::isFunctionDefinition (Clause* clause, bool inHigherOrder)
{
if (clause->length() != 1) {
return 0;
}
return isFunctionDefinition((*clause)[0],inHigherOrder);
}
FunctionDefinition::Def*
FunctionDefinition::isFunctionDefinition (Literal* lit, bool inHigherOrder)
{
if (! lit->isPositive() ||
! lit->isEquality() ||
! lit->shared()) {
return 0;
}
TermList* args = lit->args();
if (args->isVar()) {
return 0;
}
Term* l = args->term();
args = args->next();
if (args->isVar()) {
return 0;
}
Term* r = args->term();
Def* def = defines(l,r,inHigherOrder);
if (def) {
return def;
}
def = defines(r,l,inHigherOrder);
if (def) {
return def;
}
return 0;
}
FunctionDefinition::Def*
FunctionDefinition::defines (Term* lhs, Term* rhs, bool inHigherOrder)
{
if(!lhs->shared() || !rhs->shared()) {
return 0;
}
unsigned f = lhs->functor();
if(env.signature->getFunction(f)->protectedSymbol()) {
return 0;
}
if(env.signature->getFunction(f)->distinctGroups()!=0) {
return 0;
}
if(lhs->color()==COLOR_TRANSPARENT && rhs->color()!=COLOR_TRANSPARENT) {
return 0;
}
if (occurs(f,*rhs)) {
return 0;
}
if (!lhs->arity()) {
if(env.signature->isFoolConstantSymbol(true , f) ||
env.signature->isFoolConstantSymbol(false, f)){
return 0;
}
if (rhs->arity() && !inHigherOrder) { return 0;
}
if (rhs->functor() == f) {
return 0;
}
if(!inHigherOrder){
return new Def(lhs,rhs,true,true);
}
}
int vars = 0;
ZIArray<unsigned> counter;
for (const TermList* ts = lhs->args(); ts->isNonEmpty(); ts=ts->next()) {
if (! ts->isVar()) {
return 0;
}
int w = ts->var();
if (counter[w]++) { return 0;
}
vars++;
}
bool linear = true;
TermVarIterator vs(rhs->args());
while (vs.hasNext()) {
int v = vs.next();
switch (counter.get(v)) {
case 0: return 0;
case 1: counter[v]++;
vars--;
break;
default: linear = false;
break;
}
}
Def* res=new Def(lhs,rhs,linear,!vars);
if(!lhs->arity() && !rhs->arity()){
res->twoConstDef = true;
}
return res;
}
bool FunctionDefinition::occurs (unsigned f, Term& t)
{
TermFunIterator funs(&t);
while (funs.hasNext()) {
if (f == funs.next()) {
return true;
}
}
return false;
}
FunctionDefinition::Def*
FunctionDefinition::isFunctionDefinition (FormulaUnit& unit, bool inHigherOrder)
{
Formula* f = unit.formula();
while (f->connective() == FORALL) {
f = f->qarg();
}
if (f->connective() != LITERAL) {
return 0;
}
return isFunctionDefinition(f->literal(), inHigherOrder);
}
void FunctionDefinition::deleteDef (Def* def)
{
delete def;
}
void FunctionDefinitionExtra::output(std::ostream &out) const {
bool first = true;
out << "inlined=[";
for(Term *t : lhs) {
if(!first)
out << ",";
first = false;
out << t->toString();
}
out << "]";
}
}