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 }")
+ );
}
}