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.EvaluatableTerm;
import models.terms.PrimedTerm;
import models.terms.Resource;
public class RewriteInferenceSystem {
List<EquationFormula> constraintFormulas = new ArrayList<>();
List<EquationFormula> invariantFormulas = new ArrayList<>();
EquationFormula inputFormula;
List<EquationFormula> conditionalFormulas = new ArrayList<>();
List<Formula> otherFormulas = new ArrayList<>();
EquationFormula conclusion;
public RewriteInferenceSystem(List<Formula> assumptions, Formula conclusion) {
for (Formula assumption : assumptions) {
if (inputFormulaCheck(assumption)) {
inputFormula = (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 void debug() {
System.out.println("constraintFormulas: " + constraintFormulas.toString());
System.out.println("invariantFormulas: " + invariantFormulas.toString());
System.out.println("inputFormula: " + inputFormula.toString());
System.out.println("conditionalFormulas: " + conditionalFormulas.toString());
System.out.println("conclusion: " + conclusion.toString());
}
public boolean inference() {
ResourceTree baseTree = expandTree();
Set<ResourceTree> result = rewriteTree(baseTree);
for (ResourceTree tree: result) {
System.out.println("=====================================");
tree.debugAllPath();
}
return false;
}
private ResourceTree expandTree() {
ResourceTree inputResourceTree = new ResourceTree(inputFormula.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, Resource> 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 Set<ResourceTree> rewriteTree(ResourceTree expandedInputResourceTree) {
Set<ResourceTree> result = new HashSet<>();
result.add(expandedInputResourceTree);
Map<ResourceTree, ResourceTree> rewritable = new HashMap<>();
for (EquationFormula formula : constraintFormulas) {
rewritable.put(new ResourceTree(formula.getRightSideHand()), new ResourceTree(formula.getLeftSideHand()));
}
for (EquationFormula formula : conditionalFormulas) {
rewritable.put(new ResourceTree(formula.getRightSideHand()), new ResourceTree(formula.getLeftSideHand()));
}
rewritable.put(new ResourceTree(inputFormula.getLeftSideHand()), new ResourceTree(inputFormula.getRightSideHand()));
Deque<ResourceTree> treeQueue = new ArrayDeque<>();
treeQueue.add(expandedInputResourceTree);
while (! treeQueue.isEmpty()) {
ResourceTree curBaseTree = treeQueue.pollFirst();
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);
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, Resource> 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 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;
}
}