package inference.rewrite;
import java.util.ArrayDeque;
import java.util.ArrayList;
import java.util.Deque;
import java.util.HashMap;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Set;
import models.formulas.EquationFormula;
import models.formulas.Formula;
import models.formulas.Then;
import models.terms.DependencyTerm;
import models.terms.EvaluatableTerm;
import models.terms.PrimedTerm;
import models.terms.Resource;
public class RewriteInferenceSystem {
private List<EquationFormula> constraintFormulas = new ArrayList<>();
private List<EquationFormula> invariantFormulas = new ArrayList<>();
private List<EquationFormula> inputFormulas = new ArrayList<>();
private List<EquationFormula> conditionalFormulas = new ArrayList<>();
private List<Formula> otherFormulas = new ArrayList<>();
private EquationFormula conclusion;
public RewriteInferenceSystem(List<Formula> assumptions, Formula conclusion) {
for (Formula assumption : assumptions) {
if (inputFormulaCheck(assumption)) {
inputFormulas.add((EquationFormula) assumption);
} else if(invariantFormulaCheck(assumption)) {
invariantFormulas.add((EquationFormula) assumption);
} else if (assumption instanceof EquationFormula) {
constraintFormulas.add((EquationFormula) assumption);
} else {
otherFormulas.add(assumption);
}
}
if (conclusion instanceof Then then) {
conditionalFormulas.add((EquationFormula) then.getCondition());
this.conclusion = (EquationFormula) then.getResult();
} else {
this.conclusion = (EquationFormula) conclusion;
}
}
public RewriteInferenceSystem(List<EquationFormula> constraintFormulas, List<EquationFormula> invariantFormulas, List<EquationFormula> inputFormulas, Formula conclusion) {
this.constraintFormulas = constraintFormulas;
this.invariantFormulas = invariantFormulas;
this.inputFormulas = inputFormulas;
if (conclusion instanceof Then then) {
conditionalFormulas.add((EquationFormula) then.getCondition());
this.conclusion = (EquationFormula) then.getResult();
} else {
this.conclusion = (EquationFormula) conclusion;
}
}
public void debug() {
System.out.println("constraintFormulas: " + constraintFormulas.toString());
System.out.println("invariantFormulas: " + invariantFormulas.toString());
System.out.println("inputFormula: " + inputFormulas.toString());
System.out.println("conditionalFormulas: " + conditionalFormulas.toString());
System.out.println("conclusion: " + conclusion.toString());
}
public boolean inference() {
EquationFormula defCheck = definitonFormulaCheck();
if (defCheck != null) {
System.out.println("non definition formula is detected. " + defCheck.toString() );
return false;
}
if (!loopCheck()) {
System.out.println("Loop is detected");
return false;
}
Set<ResourceTree> conclusionLeftResult = new HashSet<>();
Set<ResourceTree> conclusionRightResult = new HashSet<>();
for (EquationFormula inputFormula : inputFormulas) {
// System.out.println("---------------------------" + inputFormula + "-------------------------------------------------");
// ResourceTree baseTree = expandTree(i);
// System.out.println("baseTree: " + baseTree);
// Set<ResourceTree> result = rewriteTree(baseTree, i);
// System.out.println("=================result====================");
// for (ResourceTree tree: result) {
// tree.debugAllPath();
// }
// if (result.contains(new ResourceTree(conclusion.getLeftSideHand())) && result.contains(new ResourceTree(conclusion.getRightSideHand()))) {
// System.out.println("yes");
// } else {
// System.out.println("no");
// }
// System.out.println("=====================================");
Map<ResourceTree, Set<RewriteGraphNode>> leftRewriteGraph = new HashMap<>();
Map<ResourceTree, Set<RewriteGraphNode>> rightRewriteGraph = new HashMap<>();
conclusionLeftResult.addAll(rewriteTree(new ResourceTree(conclusion.getLeftSideHand()), inputFormula, leftRewriteGraph));
conclusionRightResult.addAll(rewriteTree(new ResourceTree(conclusion.getRightSideHand()), inputFormula, rightRewriteGraph));
// System.out.println(conclusionRightResult);
// if (conclusionLeftResult.size() != 0) {
// ResourceTree resultRoot = conclusionLeftResult.iterator().next();
// showRewriteGraph(resultRoot, leftRewriteGraph);
// System.out.println("=============================================================----");
// showRewriteGraph(resultRoot, rightRewriteGraph);
// System.out.println(conclusionLeftResult);
// }
}
conclusionLeftResult.retainAll(conclusionRightResult);
// System.out.println(conclusionLeftResult.size() != 0);
return conclusionLeftResult.size() != 0;
}
private ResourceTree expandTree(int inputFormulaIndex) {
ResourceTree inputResourceTree = new ResourceTree(inputFormulas.get(inputFormulaIndex).getLeftSideHand());
List<ResourceTree> constraintResourceTree = constraintFormulas.stream().map(v -> new ResourceTree(v.getRightSideHand())).toList();
List<ResourceTree> conditionalResourceTree = conditionalFormulas.stream().map(v -> new ResourceTree(v.getLeftSideHand())).toList();
boolean rewrited = true;
while (rewrited) {
rewrited = false;
inputResourceTree.debugAllPath();
for (ResourceTree tree: constraintResourceTree) {
Position matchPos = treeJoinCheck(inputResourceTree, tree);
if (matchPos != null) {
inputResourceTree = joinTree(inputResourceTree, matchPos, tree);
rewrited = true;
} else {
matchPos = treeJoinCheck(tree, inputResourceTree);
if (matchPos != null) {
inputResourceTree = joinTree(tree, matchPos, inputResourceTree);
rewrited = true;
}
}
}
for (ResourceTree tree: conditionalResourceTree) {
Position matchPos = treeJoinCheck(inputResourceTree, tree);
if (matchPos != null) {
inputResourceTree = joinTree(inputResourceTree, matchPos, tree);
rewrited = true;
} else {
matchPos = treeJoinCheck(tree, inputResourceTree);
if (matchPos != null) {
inputResourceTree = joinTree(tree, matchPos, inputResourceTree);
rewrited = true;
}
}
}
}
return inputResourceTree;
}
private Position treeJoinCheck(ResourceTree leftTree, ResourceTree rightTree) {
Deque<Position> positionQueue = new ArrayDeque<>();
for (Position nextPos: leftTree.getChildren(new Position())) {
positionQueue.add(nextPos);
}
while(! positionQueue.isEmpty()) {
Position curPos = positionQueue.pollFirst();
if (treeJoinCheck(leftTree, curPos, rightTree, new Position())) {
return curPos;
}
for (Position nextPos: leftTree.getChildren(curPos)) {
positionQueue.add(nextPos);
}
}
return null;
}
private boolean treeJoinCheck(ResourceTree leftTree, Position leftTreePos, ResourceTree rightTree, Position rightTreePos) {
if (leftTree.getResource(leftTreePos).equals(rightTree.getResource(rightTreePos))) {
boolean result = true;
Set<Position> used = new HashSet<>();
for (Position nextLeftPos: leftTree.getChildren(leftTreePos)) {
boolean someTreeMatched = false;
for (Position nextRightPos: rightTree.getChildren(rightTreePos)) {
if (treeJoinCheck(leftTree, nextLeftPos, rightTree, nextRightPos)) {
used.add(nextRightPos);
someTreeMatched = true;
break;
}
}
result &= someTreeMatched;
}
return result;
}
return false;
}
private record PositionPair(Position resultPos, Position rightPos) {};
private ResourceTree joinTree(ResourceTree leftTree, Position joinPos, ResourceTree rightTree) {
Map<Position, List<Position>> tree = new HashMap<>();
Map<Position, EvaluatableTerm> resourceMap = new HashMap<>();
Deque<PositionPair> positionQue = new ArrayDeque<>();
Position rootPos = new Position();
positionQue.add(new PositionPair(rootPos, null));
tree.put(rootPos, new ArrayList<>());
resourceMap.put(rootPos, leftTree.getResource(rootPos));
while (! positionQue.isEmpty()) {
PositionPair curPosPair = positionQue.pollFirst();
Position curPos = curPosPair.resultPos();
Position curRightPos = curPosPair.rightPos();
if (curRightPos == null) {
for (Position nextPos: leftTree.getChildren(curPos)) {
tree.get(curPos).add(nextPos);
tree.put(nextPos, new ArrayList<>());
resourceMap.put(nextPos, leftTree.getResource(nextPos));
if (nextPos.startWith(joinPos)) {
positionQue.add(new PositionPair(nextPos, new Position()));
} else {
positionQue.add(new PositionPair(nextPos, null));
}
}
} else {
for (int i = 0; i < rightTree.getChildren(curRightPos).size(); i++) {
Position nextRightPos = rightTree.getChildren(curRightPos).get(i);
Position nextResultPos = curPos.addPath(i);
tree.get(curPos).add(nextResultPos);
tree.put(nextResultPos, new ArrayList<>());
resourceMap.put(nextResultPos, rightTree.getResource(nextRightPos));
positionQue.add(new PositionPair(nextResultPos, nextRightPos));
}
}
}
return new ResourceTree(tree, resourceMap);
}
private void showRewriteGraph(ResourceTree root, Map<ResourceTree, Set<RewriteGraphNode>> rewriteGraph) {
List<ResourceTree> result = new ArrayList<>();
List<RewriteGraphNode> resultNodes = new ArrayList<>();
result.add(root);
boolean isChanged = true;
while (isChanged) {
ResourceTree curTree = result.get(result.size() - 1);
isChanged = false;
if (rewriteGraph.containsKey(curTree)) {
isChanged |= !rewriteGraph.get(curTree).isEmpty();
RewriteGraphNode nextNode = rewriteGraph.get(curTree).iterator().next();
result.add(nextNode.baseTree());
resultNodes.add(nextNode);
}
}
for (int i = resultNodes.size() - 1; i >= 0; i--) {
System.out.println(resultNodes.get(i));
System.out.println("--------------------------");
}
System.out.println(root);
}
private record RewriteGraphNode(ResourceTree baseTree, ResourceTree leftSideHand, ResourceTree rightSideHand, ResourceTree result) {
@Override
public String toString() {
return baseTree.toString() + " --( " + leftSideHand.toString() + " = " + rightSideHand.toString() + " )--> " + result ;
}
@Override
public int hashCode() {
return toString().hashCode();
}
@Override
public boolean equals(Object another) {
if ( ! (another instanceof RewriteGraphNode)) {
return false;
}
RewriteGraphNode node = (RewriteGraphNode) another;
return this.baseTree.equals(node.baseTree()) && this.leftSideHand.equals(node.leftSideHand()) && this.rightSideHand.equals(node.rightSideHand()) && this.result.equals(node.result());
}
};
private Set<ResourceTree> rewriteTree(ResourceTree expandedInputResourceTree, EquationFormula inputFormula, Map<ResourceTree, Set<RewriteGraphNode>> rewriteGraph) {
Set<ResourceTree> result = new HashSet<>();
Map<ResourceTree, ResourceTree> rewritable = new HashMap<>();
for (EquationFormula formula : constraintFormulas) {
rewritable.put(new ResourceTree(formula.getLeftSideHand()), new ResourceTree(formula.getRightSideHand()));
}
for (EquationFormula formula : conditionalFormulas) {
rewritable.put(new ResourceTree(formula.getLeftSideHand()), new ResourceTree(formula.getRightSideHand()));
}
for (EquationFormula formula : invariantFormulas) {
rewritable.put(new ResourceTree(formula.getLeftSideHand()), new ResourceTree(formula.getRightSideHand()));
}
rewritable.put(new ResourceTree(inputFormula.getLeftSideHand()), new ResourceTree(inputFormula.getRightSideHand()));
Deque<ResourceTree> treeQueue = new ArrayDeque<>();
treeQueue.add(expandedInputResourceTree);
Set<RewriteGraphNode> used = new HashSet<>();
while (! treeQueue.isEmpty()) {
ResourceTree curBaseTree = treeQueue.pollFirst();
if (result.contains(curBaseTree)) {
continue;
}
result.add(curBaseTree);
for (ResourceTree from: rewritable.keySet()) {
ResourceTree to = rewritable.get(from);
Set<Position> matchPositions = treeMatchCheck(curBaseTree, from);
if (matchPositions == null) {
continue;
}
ResourceTree res = rewrite(curBaseTree, matchPositions, to);
if (! rewriteGraph.containsKey(res)) {
rewriteGraph.put(res, new HashSet<>());
}
RewriteGraphNode nextNode = new RewriteGraphNode(curBaseTree, from, to, res);
if (!used.contains(nextNode)) {
rewriteGraph.get(res).add(nextNode);
used.add(nextNode);
}
treeQueue.add(res);
}
}
return result;
}
private Set<Position> treeMatchCheck(ResourceTree baseTree, ResourceTree matchTree) {
Deque<Position> positionQueue = new ArrayDeque<>();
positionQueue.add(new Position());
while (! positionQueue.isEmpty()) {
Position curPos = positionQueue.pollFirst();
Set<Position> result = treeMatchCheck(baseTree, curPos, matchTree);
if (result != null) {
return result;
}
for (Position child: baseTree.getChildren(curPos)) {
positionQueue.add(child);
}
}
return null;
}
private Set<Position> treeMatchCheck(ResourceTree baseTree, Position startPosition, ResourceTree matchTree) {
Set<Position> result = new HashSet<>();
Deque<Position> basePosQueue = new ArrayDeque<>();
Deque<Position> matchPosQueue = new ArrayDeque<>();
basePosQueue.add(startPosition);
matchPosQueue.add(new Position());
if (! baseTree.getResource(startPosition).equals(matchTree.getResource(new Position()))) {
return null;
}
while (! basePosQueue.isEmpty()) {
Position curBasePos = basePosQueue.pollFirst();
Position curMatchPos = matchPosQueue.pollFirst();
result.add(curBasePos);
if (matchTree.getChildren(curMatchPos).size() == 0) {
continue;
}
Set<Position> used = new HashSet<>();
for (Position nextBasePos: baseTree.getChildren(curBasePos)) {
for (Position nextMatchPos: matchTree.getChildren(curMatchPos)) {
if (used.contains(nextMatchPos)) continue;
if (baseTree.getResource(nextBasePos).equals(matchTree.getResource(nextMatchPos))) {
used.add(nextMatchPos);
basePosQueue.add(nextBasePos);
matchPosQueue.add(nextMatchPos);
break;
}
return null;
}
}
}
return result;
}
private ResourceTree rewrite(ResourceTree baseTree, Set<Position> matchPositions, ResourceTree toTree) {
Map<Position, List<Position>> resultTree = new HashMap<>();
Map<Position, EvaluatableTerm> resultResourceMap = new HashMap<>();
Position resultCurPos = new Position();
Position baseCurPos = new Position();
Position toCurPos = new Position();
Deque<Position> resultPosQueue = new ArrayDeque<>();
Deque<Position> basePosQueue = new ArrayDeque<>();
Deque<Position> toPosQueue = new ArrayDeque<>();
resultPosQueue.add(resultCurPos);
basePosQueue.add(baseCurPos);
toPosQueue.add(toCurPos);
while (! basePosQueue.isEmpty()) {
baseCurPos = basePosQueue.pollFirst();
if (matchPositions.contains(baseCurPos)) {
List<Position> lastPositions = new ArrayList<>();
Deque<Position> matchPosQueue = new ArrayDeque<>();
matchPosQueue.add(baseCurPos);
while (! matchPosQueue.isEmpty()) {
Position matchPos = matchPosQueue.pollFirst();
for (Position nextMatchPos: baseTree.getChildren(matchPos)) {
if (matchPositions.contains(nextMatchPos)) {
matchPosQueue.add(nextMatchPos);
} else {
lastPositions.add(nextMatchPos);
}
}
}
List<Position> resultLastPositions = new ArrayList<>();
while (! toPosQueue.isEmpty()) {
toCurPos = toPosQueue.pollFirst();
resultCurPos = resultPosQueue.pollFirst();
resultTree.put(resultCurPos, new ArrayList<>());
resultResourceMap.put(resultCurPos, toTree.getResource(toCurPos));
if (toTree.getChildren(toCurPos).size() == 0) {
resultLastPositions.add(resultCurPos);
continue;
}
for (int i = 0; i < toTree.getChildren(toCurPos).size(); i++) {
Position resultNextPos = resultCurPos.addPath(i);
resultTree.get(resultCurPos).add(resultNextPos);
resultPosQueue.add(resultNextPos);
toPosQueue.add(toCurPos.addPath(i));
}
}
for (Position resultLastPos: resultLastPositions) {
Deque<Position> resultLastPosQueue = new ArrayDeque<>();
Deque<Position> lastPosQueue = new ArrayDeque<>();
for (int i = 0; i < lastPositions.size(); i++) {
Position nextResultLastPos = resultLastPos.addPath(i);
resultTree.get(resultLastPos).add(nextResultLastPos);
resultLastPosQueue.add(nextResultLastPos);
lastPosQueue.add(lastPositions.get(i));
}
while (! resultLastPosQueue.isEmpty()) {
Position curResultLastPos = resultLastPosQueue.pollFirst();
Position curLastPos = lastPosQueue.pollFirst();
resultTree.put(curResultLastPos, new ArrayList<>());
resultResourceMap.put(curResultLastPos, baseTree.getResource(curLastPos));
for (int i = 0; i < baseTree.getChildren(curLastPos).size(); i++) {
Position nextResultLastPos = curResultLastPos.addPath(i);
resultTree.get(curResultLastPos).add(nextResultLastPos);
resultLastPosQueue.add(nextResultLastPos);
lastPosQueue.add(curLastPos.addPath(i));
}
}
}
} else {
resultCurPos = resultPosQueue.pollFirst();
resultTree.put(resultCurPos, new ArrayList<>());
resultResourceMap.put(resultCurPos, baseTree.getResource(baseCurPos));
for (int i = 0; i < baseTree.getChildren(baseCurPos).size(); i++) {
Position resultNextPos = resultCurPos.addPath(i);
resultTree.get(resultCurPos).add(resultNextPos);
resultPosQueue.add(resultNextPos);
basePosQueue.add(baseCurPos.addPath(i));
}
}
}
return new ResourceTree(resultTree, resultResourceMap);
}
private boolean loopCheck() {
Map<EvaluatableTerm, EvaluatableTerm> equationGraph = new HashMap<>();
for (EquationFormula formula : constraintFormulas) {
equationGraph.put(formula.getLeftSideHand(), formula.getRightSideHand());
}
for (EquationFormula formula : invariantFormulas) {
equationGraph.put(formula.getLeftSideHand(), formula.getRightSideHand());
}
for (EquationFormula inputFormula : inputFormulas) {
equationGraph.put(inputFormula.getLeftSideHand(), inputFormula.getRightSideHand());
}
for (EquationFormula formula : conditionalFormulas) {
equationGraph.put(formula.getLeftSideHand(), formula.getRightSideHand());
}
for (EvaluatableTerm leftTerm : equationGraph.keySet()) {
Set<EvaluatableTerm> used = new HashSet<>();
used.add(leftTerm);
EvaluatableTerm rightTerm = equationGraph.get(leftTerm);
if (checkUsedTerms(used, rightTerm)) {
return false;
}
while (equationGraph.containsKey(rightTerm)) {
leftTerm = rightTerm;
used.add(leftTerm);
rightTerm = equationGraph.get(leftTerm);
if (checkUsedTerms(used, rightTerm)) {
return false;
}
}
}
return true;
}
private boolean checkUsedTerms(Set<EvaluatableTerm> used, EvaluatableTerm term) {
if (used.contains(term)) {
return true;
}
if (term instanceof DependencyTerm depTerm) {
if (checkUsedTerms(used, depTerm.getDependingTerm())) {
return true;
}
for (EvaluatableTerm te : depTerm.getDependedTerms()) {
if (checkUsedTerms(used, te)) {
return true;
}
}
for (EvaluatableTerm te : depTerm.getArgumentTerms()) {
if (checkUsedTerms(used, te)) {
return true;
}
}
}
return false;
}
private EquationFormula definitonFormulaCheck() {
for (EquationFormula eq : constraintFormulas) {
if (!definitionFormulaCheck(eq)) {
return eq;
}
}
for (EquationFormula inputFormula : inputFormulas) {
if (!definitionFormulaCheck(inputFormula)) {
return inputFormula;
}
}
for (EquationFormula eq : invariantFormulas) {
if (!definitionFormulaCheck(eq)) {
return eq;
}
}
for (EquationFormula eq : conditionalFormulas) {
if (!definitionFormulaCheck(eq)) {
return eq;
}
}
return null;
}
private boolean definitionFormulaCheck(Formula formula) {
if (formula instanceof EquationFormula equation) {
return singleTermCheck(equation.getLeftSideHand());
} else if (formula instanceof Then then) {
Formula f1 = then.getCondition();
Formula f2 = then.getResult();
if (f1 instanceof EquationFormula eq1 && f2 instanceof EquationFormula eq2) {
return singleTermCheck(eq1.getLeftSideHand()) && singleTermCheck(eq2.getLeftSideHand());
}
}
return false;
}
private boolean singleTermCheck(EvaluatableTerm term) {
if (term instanceof PrimedTerm primedTerm) {
return singleTermCheck(primedTerm.getPrimedTerm());
}
if (term instanceof Resource) {
return true;
}
DependencyTerm depTerm = (DependencyTerm) term;
if (!(depTerm.getDependingTerm() instanceof Resource)) {
if (depTerm.getDependingTerm() instanceof PrimedTerm primedTerm) {
if (! (primedTerm.getPrimedTerm() instanceof Resource)) {
return false;
}
} else {
return false;
}
}
for (EvaluatableTerm te : depTerm.getDependedTerms()) {
if (!(te instanceof Resource)) {
if (te instanceof PrimedTerm primedTerm) {
if (! (primedTerm.getPrimedTerm() instanceof Resource)) {
return false;
}
} else {
return false;
}
}
}
for (EvaluatableTerm te : depTerm.getArgumentTerms()) {
if (!(te instanceof Resource)) {
if (te instanceof PrimedTerm primedTerm) {
if (! (primedTerm.getPrimedTerm() instanceof Resource)) {
return false;
}
} else {
return false;
}
}
}
return true;
}
private boolean inputFormulaCheck(Formula formula) {
if (! (formula instanceof EquationFormula equation)) {
return false;
}
List<Resource> resources = new ArrayList<>(equation.getLeftSideHand().getSubTerms(Resource.class).values());
resources.addAll(equation.getRightSideHand().getSubTerms(Resource.class).values());
for (Resource resource : resources) {
if (resource.getOrder() == 0) {
return true;
}
}
return false;
}
private boolean invariantFormulaCheck(Formula formula) {
if (! (formula instanceof EquationFormula equation)) {
return false;
}
EvaluatableTerm leftSideHand = equation.getLeftSideHand();
EvaluatableTerm rightSideHand = equation.getRightSideHand();
if (leftSideHand instanceof PrimedTerm) {
return ((PrimedTerm) leftSideHand).getPrimedTerm().equals(rightSideHand);
} else if (rightSideHand instanceof PrimedTerm) {
return ((PrimedTerm) rightSideHand).getPrimedTerm().equals(leftSideHand);
}
return false;
}
}