diff --git a/reactor/Core/src/jetbrains/mps/unification/TermGraphUnifier.kt b/reactor/Core/src/jetbrains/mps/unification/TermGraphUnifier.kt index 28ba2516..140927ec 100644 --- a/reactor/Core/src/jetbrains/mps/unification/TermGraphUnifier.kt +++ b/reactor/Core/src/jetbrains/mps/unification/TermGraphUnifier.kt @@ -18,6 +18,7 @@ package jetbrains.mps.unification import gnu.trove.list.array.TIntArrayList import gnu.trove.map.hash.TIntObjectHashMap +import jetbrains.mps.logic.reactor.logical.Logical import jetbrains.mps.unification.Substitution.FailureCause.CYCLE_DETECTED import jetbrains.mps.unification.Substitution.FailureCause.SYMBOL_CLASH import jetbrains.mps.unification.Term.Kind.* @@ -294,13 +295,15 @@ class TermGraphUnifier { return wrapper.unwrap(origin[t]) } - private fun idSymbol(symbol: Any): Any { - return if (symbol is String) symbol.intern() else symbol - } - + private fun idSymbol(symbol: Any): Any = + when (symbol) { + is String -> symbol.intern() + is Logical<*> -> symbol.findRoot() + else -> symbol + } + private fun Any?.eq (that: Any?): Boolean { return if (this == null) that == null else this.equals(that) } - } \ No newline at end of file diff --git a/reactor/Test/test/jetbrains/mps/unification/test/SolverTests.java b/reactor/Test/test/jetbrains/mps/unification/test/SolverTests.java index eed9f766..e273c0a6 100644 --- a/reactor/Test/test/jetbrains/mps/unification/test/SolverTests.java +++ b/reactor/Test/test/jetbrains/mps/unification/test/SolverTests.java @@ -16,6 +16,9 @@ package jetbrains.mps.unification.test; +import jetbrains.mps.logic.reactor.core.LogicalImpl; +import jetbrains.mps.logic.reactor.logical.MetaLogical; +import jetbrains.mps.unification.Substitution; import jetbrains.mps.unification.Term; import jetbrains.mps.unification.TermWrapper; import org.jetbrains.annotations.NotNull; @@ -543,4 +546,44 @@ public class SolverTests { } + @Test + public void joinedLogicals() throws Exception { + MetaLogical X = new MetaLogical<>("X", Term.class); + MetaLogical Y = new MetaLogical<>("Y", Term.class); + MetaLogical Z = new MetaLogical<>("Z", Term.class); + LogicalImpl xLogical = new LogicalImpl<>(X); + LogicalImpl yLogical = new LogicalImpl<>(Y); + LogicalImpl zLogical = new LogicalImpl<>(Z); + + Term left = term("foo", term("bar", logicalVar(yLogical)), logicalVar(zLogical)); + Term right = term("foo", term("bar", logicalVar(xLogical)), logicalVar(zLogical)); + + assertUnifiesWithBindings(left, right, + new Substitution.Binding(logicalVar(xLogical), logicalVar(yLogical))); + + xLogical.union(yLogical); + + assertUnifiesWithBindings(left, right); + } + + @Test + public void joinedLogicals_cycle() throws Exception { + MetaLogical X = new MetaLogical<>("X", Term.class); + MetaLogical Y = new MetaLogical<>("Y", Term.class); + MetaLogical Z = new MetaLogical<>("Z", Term.class); + LogicalImpl xLogical = new LogicalImpl<>(X); + LogicalImpl yLogical = new LogicalImpl<>(Y); + LogicalImpl zLogical = new LogicalImpl<>(Z); + + Term left = logicalVar(yLogical); + Term right = term("foo", term("bar", logicalVar(xLogical)), logicalVar(zLogical)); + + assertUnifiesWithBindings(left, right, + new Substitution.Binding(left, right)); + + xLogical.union(yLogical); + + assertUnificationFails(left, right, + CYCLE_DETECTED); + } }