diff --git a/reactor/code/src/jetbrains/mps/unification/UnionFindTermGraphUnifier.java b/reactor/code/src/jetbrains/mps/unification/UnionFindTermGraphUnifier.java index 5213efa0..af2d7231 100644 --- a/reactor/code/src/jetbrains/mps/unification/UnionFindTermGraphUnifier.java +++ b/reactor/code/src/jetbrains/mps/unification/UnionFindTermGraphUnifier.java @@ -16,8 +16,13 @@ package jetbrains.mps.unification; +import jetbrains.mps.unification.Unification.SuccessfulSubstitution; + import java.util.*; +import static jetbrains.mps.unification.Node.Kind.*; +import static jetbrains.mps.unification.Unification.*; + /** * This is an implementation of the "near linear" algorithm for solving syntactic unification * as described in the paper linked below and also in the textbook of the same author.1 2 @@ -44,9 +49,11 @@ public class UnionFindTermGraphUnifier { private Map myData = new IdentityHashMap(); + private int myUnreconciledRefs = 0; + public Substitution unify(Node a, Node b) { if (!unifClosure(a, b)) { - return Unification.FAILED_SUBSTITUTION; + return FAILED_SUBSTITUTION; } return findSolution(a); @@ -62,16 +69,16 @@ public class UnionFindTermGraphUnifier { Node zt = getSchema(t); // a VAR always matches another node - if(zs.is(Node.Kind.VAR) || zt.is(Node.Kind.VAR)) { + if(zs.is(VAR) || zt.is(VAR)) { union(s, t); return true; } // dereference REF nodes - Node ds = zs.is(Node.Kind.REF) ? zs.get() : zs; - Node dt = zt.is(Node.Kind.REF) ? zt.get() : zt; + Node ds = zs.is(REF) ? zs.get() : zs; + Node dt = zt.is(REF) ? zt.get() : zt; - if (ds.is(Node.Kind.VAR)) { + if (ds.is(VAR)) { if (s != find(ds)) { union(s, find(ds)); } @@ -82,7 +89,7 @@ public class UnionFindTermGraphUnifier { zs = ds; } - if (dt.is(Node.Kind.VAR)) { + if (dt.is(VAR)) { if (t != find(dt)) { union(t, find(dt)); } @@ -98,14 +105,14 @@ public class UnionFindTermGraphUnifier { return true; } - if (zs.is(Node.Kind.FUN) && zt.is(Node.Kind.FUN)) + if (zs.is(FUN) && zt.is(FUN)) { if (!eq(zs.symbol(), zt.symbol())) { return false; // symbol clash } // union REF nodes only to each other - if (s.is(Node.Kind.REF) == t.is(Node.Kind.REF)) { + if (s.is(REF) == t.is(REF)) { union(s, t); } @@ -132,7 +139,7 @@ public class UnionFindTermGraphUnifier { if (ssize < tsize) { Node tmp = t; t = s; s = tmp; } - else if (ssize == tsize && s.is(Node.Kind.VAR) && t.is(Node.Kind.VAR)) { + else if (ssize == tsize && s.is(VAR) && t.is(VAR)) { // ensure proper order of variables in the substitution if(t.compareTo(s) < 0) { Node tmp = t; t = s; s = tmp; @@ -146,8 +153,7 @@ public class UnionFindTermGraphUnifier { Node zs = getSchema(s); Node zt = getSchema(t); - if (zs.is(Node.Kind.VAR) || - (zs.is(Node.Kind.REF) && zs.get().is(Node.Kind.VAR) && !zt.is(Node.Kind.VAR))) + if (zs.is(VAR) || (zs.is(REF) && zs.get().is(VAR) && !zt.is(VAR))) { setSchema(s, zt); } @@ -177,7 +183,8 @@ public class UnionFindTermGraphUnifier { } private Substitution findSolution(Node s) { - return findSolution(s, Unification.EMPTY_SUBSTITUTION); + myUnreconciledRefs = 0; + return findSolution(s, EMPTY_SUBSTITUTION); } private Substitution findSolution(Node s, Substitution substitution) { @@ -187,10 +194,11 @@ public class UnionFindTermGraphUnifier { return substitution; // not part of a cycle } if (isVisited(z)) { - return Unification.FAILED_SUBSTITUTION; // there exists a cycle + return FAILED_SUBSTITUTION; // there exists a cycle } - if (z.is(Node.Kind.FUN)) { + int unreconciled = myUnreconciledRefs; + if (z.is(FUN)) { setVisited(z, true); for (Node c : z.children()) { @@ -208,15 +216,24 @@ public class UnionFindTermGraphUnifier { return substitution; } + if (isReferenced(z)) { + setReferenced(z, false); + myUnreconciledRefs--; + } + setAcyclic(z, true); - Unification.SuccessfulSubstitution success = new Unification.SuccessfulSubstitution(substitution); + SuccessfulSubstitution success = new SuccessfulSubstitution(substitution); for (Node var : getVars(find(z))) { if (var != z) { - Node val = z.is(Node.Kind.REF) ? z.get() : z; + if (myUnreconciledRefs != unreconciled) { + return FAILED_SUBSTITUTION; // there's an unreconciled outward reference + } + + Node val = z.is(REF) ? z.get() : z; // Keep the order of variables within a binding - if (val.is(Node.Kind.VAR) && val.compareTo(var) < 0) { + if (val.is(VAR) && val.compareTo(var) < 0) { success.addBinding(val, var); } else { @@ -225,6 +242,11 @@ public class UnionFindTermGraphUnifier { } } + if (z.is(REF) && z.get().is(FUN)) { + setReferenced(z.get(), true); + myUnreconciledRefs++; + } + return success; } @@ -256,10 +278,7 @@ public class UnionFindTermGraphUnifier { } private List getVars(Node n) { - if (!hasData(n)) { - return collectVars(n); - } - return getData(n).myVars; + return hasData(n) ? getData(n).myVars : singletonVar(n); } private void appendVars(Node n, List vars) { @@ -269,8 +288,7 @@ public class UnionFindTermGraphUnifier { } private boolean isAcyclic(Node n) { - if (!hasData(n)) return false; - return getData(n).myAcyclic; + return hasData(n) && getData(n).myAcyclic; } private void setAcyclic(Node n, boolean acyclic) { @@ -278,50 +296,58 @@ public class UnionFindTermGraphUnifier { } private boolean isVisited(Node n) { - if (!hasData(n)) return false; - return getData(n).myVisited; + return hasData(n) && getData(n).myVisited; } private void setVisited(Node n, boolean visited) { getData(n).myVisited = visited; } + private boolean isReferenced(Node n) { + return hasData(n) && getData(n).myReferenced; + } + + private void setReferenced(Node n, boolean referenced) { + getData(n).myReferenced = referenced; + } + private boolean hasData(Node n) { - Object key = n.is(Node.Kind.VAR) ? String.valueOf(n.symbol()).intern() : n; + Object key = n.is(VAR) ? String.valueOf(n.symbol()).intern() : n; return myData.containsKey(key); } private Data getData(Node n) { - Object key = n.is(Node.Kind.VAR) ? String.valueOf(n.symbol()).intern() : n; + Object key = n.is(VAR) ? String.valueOf(n.symbol()).intern() : n; if (myData.containsKey(key)) return myData.get(key); Data data = new Data(n); myData.put(key, data); return data; } - private List collectVars(Node n) { - if (n.is(Node.Kind.VAR)) { - return Collections.singletonList(n); - } - return Collections.emptyList(); - } - - private boolean eq(Object a, Object b) { + private static boolean eq(Object a, Object b) { return a == null ? b == null : a.equals(b); } - private class Data { + private static List singletonVar(Node n) { + return n.is(VAR) ? + Collections.singletonList(n) : + Collections.emptyList(); + } + + private static class Data { int mySize = 1; boolean myAcyclic = false; boolean myVisited = false; - List myVars; + boolean myReferenced = false; + List myVars; Node myClass; Node mySchema; + Data(Node n) { myClass = n; mySchema = n; - myVars = collectVars(n); + myVars = singletonVar(n); } } } diff --git a/reactor/tests/src/jetbrains/mps/unification/test/AssertStructurallyEquivalent.java b/reactor/tests/src/jetbrains/mps/unification/test/AssertStructurallyEquivalent.java index 1ff68c20..16f3d3b1 100644 --- a/reactor/tests/src/jetbrains/mps/unification/test/AssertStructurallyEquivalent.java +++ b/reactor/tests/src/jetbrains/mps/unification/test/AssertStructurallyEquivalent.java @@ -84,10 +84,11 @@ public class AssertStructurallyEquivalent { private static class Signature { + private NodeWalker[] walkers; + private IdentityHashMap labels = new IdentityHashMap(); private int label = 1; private StringBuilder signature = new StringBuilder(); - private NodeWalker[] walkers; protected void label(Node node) { labels.put(node, label++); diff --git a/reactor/tests/src/jetbrains/mps/unification/test/SolverTests.java b/reactor/tests/src/jetbrains/mps/unification/test/SolverTests.java index 094b2dc9..06372b73 100644 --- a/reactor/tests/src/jetbrains/mps/unification/test/SolverTests.java +++ b/reactor/tests/src/jetbrains/mps/unification/test/SolverTests.java @@ -382,31 +382,15 @@ public class SolverTests { } @Test - public void testFail1() throws Exception { + public void testFailConflict() throws Exception { assertUnificationFails( term("a"), term("b") ); - } - - @Test - public void testFail2() throws Exception { - assertUnificationFails( - parse("a{b c}"), - parse("a{X}") - ); - } - - @Test - public void testFail3() throws Exception { assertUnificationFails( parse("node{name{X} child{abc}}"), parse("node{name{foo} child{X}}") ); - } - - @Test - public void testFail4() throws Exception { assertUnificationFails( parse("f{a{X} Y }"), parse("f{Y a{b{X}}}") @@ -414,18 +398,31 @@ public class SolverTests { } @Test - public void testFail5() throws Exception { + public void testFailCard() throws Exception { assertUnificationFails( - parse("f{X}"), - parse("X") + parse("a{b c}"), + parse("a{X}") ); } @Test - public void testFail6() throws Exception { + public void testFailRecursive() throws Exception { + assertUnificationFails( + parse("f{X}"), + parse("X") + ); assertUnificationFails( parse("f{f{X}}"), parse("f{X}") ); + assertUnificationFails( + parse("a{X c{X}}"), + parse("a{b{Y} Y }") + + ); + assertUnificationFails ( + parse("a{@1 b{c{^1}} @2 c{b{^2}}}"), + parse("a{ b{Y} Y }") + ); } }