diff --git a/src/main/java/Main.java b/src/main/java/Main.java index 89cd088..1058c2e 100644 --- a/src/main/java/Main.java +++ b/src/main/java/Main.java @@ -2,11 +2,14 @@ import inference.rewrite.Position; import inference.rewrite.ResourceTree; +import inference.rewrite.RewriteInferenceSystem; import lombok.SneakyThrows; import models.algebra.Expression; import models.algebra.Type; import models.dataConstraintModel.DataConstraintModel; import models.dataFlowModel.DataTransferModel; +import models.formulas.EquationFormula; +import models.formulas.Then; import models.terms.DependencyTerm; import models.terms.Resource; import parser.Parser; @@ -21,7 +24,11 @@ static Type INT = DataConstraintModel.typeInt; public static void main(String[] args) { - sandbox2(); + sandbox4(); + System.out.println("====================================="); + System.out.println("====================================="); + System.out.println("====================================="); + sandbox5(); } static void sandbox1() { @@ -71,6 +78,67 @@ } + static void sandbox3() { + Resource A = new Resource("A", INT, 1); + Resource B = new Resource("B", INT, 1); + Resource C = new Resource("C", INT, 1); + Resource D = new Resource("D", INT, 1); + Resource x = new Resource("x", INT, 0); + Resource y = new Resource("y", INT, 1); + DependencyTerm t1 = new DependencyTerm(A, B, C); + DependencyTerm t2 = new DependencyTerm(B, C, D); + EquationFormula f1 = new EquationFormula(t1, x); + EquationFormula f2 = new EquationFormula(y, t2); + RewriteInferenceSystem ris = new RewriteInferenceSystem(List.of(f1, f2), null); + ris.inference(); + + } + + static void sandbox4() { + Resource uadd = new Resource("uadd", INT, 1); + Resource add = new Resource("add", INT, 1); + Resource cid = new Resource("cid", INT, 1); + Resource org = new Resource("org", INT, 1); + Resource uid = new Resource("uid", INT, 1); + Resource x = new Resource("x", INT, 0); + Resource y = new Resource("y", INT, 0); + DependencyTerm t1 = new DependencyTerm(add, cid, org); + DependencyTerm t2 = new DependencyTerm(org, uid, x); + DependencyTerm t3 = new DependencyTerm(uadd, uid, x); + DependencyTerm t4 = new DependencyTerm(add, cid, y); + EquationFormula f1 = new EquationFormula(uadd, t1); + EquationFormula f2 = new EquationFormula(t2, y); + EquationFormula f3 = new EquationFormula(t3, t4); + RewriteInferenceSystem ris = new RewriteInferenceSystem(List.of(f1, f2), f3); + ris.inference(); + + } + + static void sandbox5() { + Resource uadd = new Resource("uadd", INT, 1); + Resource add = new Resource("add", INT, 1); + Resource cid = new Resource("cid", INT, 1); + Resource org = new Resource("org", INT, 1); + Resource uid = new Resource("uid", INT, 1); + Resource x = new Resource("x", INT, 0); + Resource y = new Resource("y", INT, 0); + Resource z = new Resource("z", INT, 0); + DependencyTerm t1 = new DependencyTerm(add, cid, org); + DependencyTerm t2 = new DependencyTerm(add, cid, x); + DependencyTerm t3 = new DependencyTerm(org, uid, z); + DependencyTerm t4 = new DependencyTerm(uadd, uid, z); + + EquationFormula f1 = new EquationFormula(uadd, t1); + EquationFormula f2 = new EquationFormula(t2, y); + EquationFormula f3 = new EquationFormula(t3, x); + EquationFormula f4 = new EquationFormula(t4, y); + Then f5 = new Then(f3, f4); + + RewriteInferenceSystem ris = new RewriteInferenceSystem(List.of(f1, f2), f5); + ris.inference(); + + } + static void sandbox7() { Resource A = new Resource("A", INT, 1); diff --git a/src/main/java/inference/rewrite/Position.java b/src/main/java/inference/rewrite/Position.java index f8dc549..61755b5 100644 --- a/src/main/java/inference/rewrite/Position.java +++ b/src/main/java/inference/rewrite/Position.java @@ -25,6 +25,27 @@ return new Position(Collections.unmodifiableList(nextPaths)); } + public boolean startWith(Position pos) { + for (int i = 0; i < pos.size(); i++) { + if (getPath(i) != pos.getPath(i)) { + return false; + } + } + return true; + } + + public int size() { + return paths.size(); + } + + private int getPath(int index) { + if (index >= size()) { + return -1; + } + return paths.get(index); + } + + @Override public boolean equals(Object another) { if (! (another instanceof Position)) { diff --git a/src/main/java/inference/rewrite/ResourceTree.java b/src/main/java/inference/rewrite/ResourceTree.java index 881a778..f379f04 100644 --- a/src/main/java/inference/rewrite/ResourceTree.java +++ b/src/main/java/inference/rewrite/ResourceTree.java @@ -4,6 +4,7 @@ import java.util.HashMap; import java.util.List; import java.util.Map; +import java.util.Objects; import java.util.stream.Collectors; import lombok.Getter; @@ -25,9 +26,35 @@ root = resourceMap.get(new Position(List.of(0))); } + public ResourceTree(Map> tree, Map resourceMap) { + this.tree = tree; + this.resourceMap = resourceMap; + root = resourceMap.get(new Position(List.of(0))); + } + + + public Resource getResource(Position pos) { + if (pos == null) { + return null; + } + return resourceMap.get(pos); + } + + public List getChildren(Position pos) { + if (tree.get(pos) == null) { + return List.of(); + } + return tree.get(pos); + } + + + private List constructResourceTree(EvaluatableTerm term, Position top) { if (term instanceof Resource resource) { resourceMap.put(top, resource); + if (! tree.containsKey(top)) { + tree.put(top, new ArrayList<>()); + } return List.of(top); } else if (term instanceof DependencyTerm depTerm) { EvaluatableTerm dependingTerm = depTerm.getDependingTerm(); @@ -58,7 +85,9 @@ @Override public String toString() { - return tree.toString(); + List result = new ArrayList<>(); + toStringAllPath(new Position(), new ArrayList<>(), result); + return result.stream().collect(Collectors.joining("\n")); } public void debug(Position pos) { @@ -86,5 +115,31 @@ curPath.remove(curPath.size() - 1); } + private void toStringAllPath(Position pos, List curPath, List result) { + curPath.add(resourceMap.get(pos)); + if (tree.get(pos).size() == 0) { + result.add(curPath.stream().map(Resource::toString).collect(Collectors.joining("-"))); + } else { + for (Position nextPos: tree.get(pos)) { + toStringAllPath(nextPos, curPath, result); + } + } + curPath.remove(curPath.size() - 1); + } + + + @Override + public boolean equals(Object another) { + if (another instanceof ResourceTree tree) { + return this.tree.equals(tree.tree) && this.resourceMap.equals(tree.resourceMap); + } + return false; + } + + @Override + public int hashCode() { + return Objects.hash(this.tree, this.resourceMap); + } + } diff --git a/src/main/java/inference/rewrite/RewriteInferenceSystem.java b/src/main/java/inference/rewrite/RewriteInferenceSystem.java index f9ea440..665bb44 100644 --- a/src/main/java/inference/rewrite/RewriteInferenceSystem.java +++ b/src/main/java/inference/rewrite/RewriteInferenceSystem.java @@ -3,26 +3,27 @@ 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 inference.rewrite.ResourceTree; 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 { - List constraintFormulas = new ArrayList<>(); - List invariantFormulas = new ArrayList<>(); + List constraintFormulas = new ArrayList<>(); + List invariantFormulas = new ArrayList<>(); EquationFormula inputFormula; - List conditionalFormulas = new ArrayList<>(); + List conditionalFormulas = new ArrayList<>(); List otherFormulas = new ArrayList<>(); - Formula conclusion; + EquationFormula conclusion; public RewriteInferenceSystem(List assumptions, Formula conclusion) { @@ -30,19 +31,19 @@ if (inputFormulaCheck(assumption)) { inputFormula = (EquationFormula) assumption; } else if(invariantFormulaCheck(assumption)) { - invariantFormulas.add(assumption); + invariantFormulas.add((EquationFormula) assumption); } else if (assumption instanceof EquationFormula) { - constraintFormulas.add(assumption); + constraintFormulas.add((EquationFormula) assumption); } else { otherFormulas.add(assumption); } } if (conclusion instanceof Then then) { - conditionalFormulas.add(then.getCondition()); - this.conclusion = then.getResult(); + conditionalFormulas.add((EquationFormula) then.getCondition()); + this.conclusion = (EquationFormula) then.getResult(); } else { - this.conclusion = conclusion; + this.conclusion = (EquationFormula) conclusion; } } @@ -56,101 +57,301 @@ } public boolean inference() { -// Map> inputFormulaTree = new HashMap<>(); -// Resource root = parseDependencyTerm(inputFormula.getLeftSideHand(), inputFormulaTree, new ArrayList<>()); -// ResourceTree inputResourceTree = new ResourceTree(root, inputFormulaTree); -// List otherRoots = new ArrayList<>(); -// for (Formula formula : constraintFormulas) { -// if (formula instanceof EquationFormula equation) { -// Map> resourceTree = new HashMap<>(); -// Resource treeRoot = parseDependencyTerm(equation.getRightSideHand(), resourceTree, new ArrayList<>()); -// otherRoots.add(new ResourceTree(treeRoot, resourceTree)); -// } -// } -// for (Formula formula : conditionalFormulas) { -// if (formula instanceof EquationFormula equation) { -// Map> resourceTree = new HashMap<>(); -// Resource treeRoot = parseDependencyTerm(equation.getLeftSideHand(), resourceTree, new ArrayList<>()); -// otherRoots.add(new ResourceTree(treeRoot, resourceTree)); -// } -// } -// System.out.println(inputResourceTree); + ResourceTree baseTree = expandTree(); + Set result = rewriteTree(baseTree); + for (ResourceTree tree: result) { + System.out.println("====================================="); + tree.debugAllPath(); + } return false; } - private Resource parseDependencyTerm(EvaluatableTerm term, Map> resourceTree, List tops) { - if (term instanceof Resource resource) { - if (tops.isEmpty()) { - resourceTree.put(resource, new ArrayList<>()); - tops.add(resource); - } else { - for (Resource top : tops) { - resourceTree.get(top).add(resource); - resourceTree.put(resource, new ArrayList<>()); + + private ResourceTree expandTree() { + ResourceTree inputResourceTree = new ResourceTree(inputFormula.getLeftSideHand()); + List constraintResourceTree = constraintFormulas.stream().map(v -> new ResourceTree(v.getRightSideHand())).toList(); + List 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; + } } - tops.clear(); - tops.add(resource); } - return resource; + 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; + } + } + } } - Resource root; - DependencyTerm depTerm = (DependencyTerm) term; - EvaluatableTerm dependingTerm = depTerm.getDependingTerm(); - List dependedTerms = depTerm.getDependedResources(); - List argumentTerms = depTerm.getArgumentTerms(); - root = parseDependencyTerm(dependingTerm, resourceTree, tops); - List nextTops = new ArrayList<>(); - for (int i = 0; i < dependedTerms.size(); i++) { - List currentTops = new ArrayList<>(tops); - parseDependencyTerm(dependedTerms.get(i), resourceTree, currentTops); - parseDependencyTerm(argumentTerms.get(i), resourceTree, currentTops); - nextTops.addAll(currentTops); - } - tops.clear(); - tops.addAll(nextTops); - return root; + return inputResourceTree; } -// private Map> joinTree(Map> baseTree, Map> joinTree) { -// -// } -// - private ResourceTree joinTreeLeft(ResourceTree baseTree, ResourceTree joinTree) { - -// ResourceTree result = new ResourceTree(baseTree.root(), new HashMap<>(baseTree.tree())); -// -// Deque curRootQue = new ArrayDeque<>(); -// curRootQue.add(joinTree.root()); -// -// while (! curRootQue.isEmpty()) { -// Resource curRoot = curRootQue.pollFirst(); -// if (isTreeMatch(result.tree(), curRoot, joinTree.tree())) { -// -// } -// curRootQue.addAll(joinTree.tree.get(curRoot)); -// } - -// return result; + private Position treeJoinCheck(ResourceTree leftTree, ResourceTree rightTree) { + Deque 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 isTreeMatch(Map> baseTree, Resource joinRoot, Map> joinTree) { - Deque joinCurResourceStack = new ArrayDeque<>(); - joinCurResourceStack.add(joinRoot); - while (! joinCurResourceStack.isEmpty()) { - Resource joinCurResource = joinCurResourceStack.pollLast(); - if (! baseTree.keySet().contains(joinCurResource)) { - return false; + private boolean treeJoinCheck(ResourceTree leftTree, Position leftTreePos, ResourceTree rightTree, Position rightTreePos) { + if (leftTree.getResource(leftTreePos).equals(rightTree.getResource(rightTreePos))) { + boolean result = true; + Set 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; } - if (! baseTree.get(joinCurResource).equals(joinTree.get(joinCurResource))) { - return false; + return result; + } + return false; + } + + private record PositionPair(Position resultPos, Position rightPos) {}; + + private ResourceTree joinTree(ResourceTree leftTree, Position joinPos, ResourceTree rightTree) { + Map> tree = new HashMap<>(); + Map resourceMap = new HashMap<>(); + + Deque 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)); + } } - joinCurResourceStack.addAll(joinTree.get(joinCurResource)); } - return true; + return new ResourceTree(tree, resourceMap); + } + + private Set rewriteTree(ResourceTree expandedInputResourceTree) { + Set result = new HashSet<>(); + result.add(expandedInputResourceTree); + Map 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 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 matchPositions = treeMatchCheck(curBaseTree, from); + if (matchPositions == null) { + continue; + } + ResourceTree res = rewrite(curBaseTree, matchPositions, to); + treeQueue.add(res); + } + } + + return result; + } + + private Set treeMatchCheck(ResourceTree baseTree, ResourceTree matchTree) { + Deque positionQueue = new ArrayDeque<>(); + positionQueue.add(new Position()); + while (! positionQueue.isEmpty()) { + Position curPos = positionQueue.pollFirst(); + Set result = treeMatchCheck(baseTree, curPos, matchTree); + if (result != null) { + return result; + } + for (Position child: baseTree.getChildren(curPos)) { + positionQueue.add(child); + } + } + return null; + } + + private Set treeMatchCheck(ResourceTree baseTree, Position startPosition, ResourceTree matchTree) { + Set result = new HashSet<>(); + Deque basePosQueue = new ArrayDeque<>(); + Deque 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 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 matchPositions, ResourceTree toTree) { + Map> resultTree = new HashMap<>(); + Map resultResourceMap = new HashMap<>(); + Position resultCurPos = new Position(); + Position baseCurPos = new Position(); + Position toCurPos = new Position(); + Deque resultPosQueue = new ArrayDeque<>(); + Deque basePosQueue = new ArrayDeque<>(); + Deque toPosQueue = new ArrayDeque<>(); + + resultPosQueue.add(resultCurPos); + basePosQueue.add(baseCurPos); + toPosQueue.add(toCurPos); + + + while (! basePosQueue.isEmpty()) { + baseCurPos = basePosQueue.pollFirst(); + if (matchPositions.contains(baseCurPos)) { + List lastPositions = new ArrayList<>(); + Deque 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 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 resultLastPosQueue = new ArrayDeque<>(); + Deque 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;