diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/core/TermTrie.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/core/TermTrie.kt index 9566180a..63d3b4f4 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/TermTrie.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/TermTrie.kt @@ -4,10 +4,7 @@ import com.github.andrewoma.dexx.collection.ConsList import com.github.andrewoma.dexx.collection.ConsList.empty import com.github.andrewoma.dexx.collection.Map as PersMap import com.github.andrewoma.dexx.collection.Maps -import jetbrains.mps.logic.reactor.util.IdHashSet -import jetbrains.mps.logic.reactor.util.cons -import jetbrains.mps.logic.reactor.util.emptySet -import jetbrains.mps.logic.reactor.util.remove +import jetbrains.mps.logic.reactor.util.* import jetbrains.mps.unification.Term import java.util.* @@ -54,13 +51,13 @@ class TermTrie() { this.root = setRoot } - fun put(term: Term, value: T): TermTrie = TermTrie(putValue(root, value, term, emptyList())) + fun put(term: Term, value: T): TermTrie = TermTrie(putValue(root, value, IdHashSet(), term, empty())) - fun remove(term: Term, value: T): TermTrie = TermTrie(removeValue(root, value, term, emptyList())) + fun remove(term: Term, value: T): TermTrie = TermTrie(removeValue(root, value, IdHashSet(), term, empty())) fun lookupValues(term: Term): Iterable { val result = ArrayList() - visitMatching(root, term, emptyList()) { value -> result.add(value) } + visitMatching(root, IdHashSet(), term, empty()) { value -> result.add(value) } return result } @@ -70,10 +67,10 @@ class TermTrie() { return result } - private fun putValue(node: PathNode, value: T, term: Term, tail: List): PathNode { - val deref = deref(term) + private fun putValue(node: PathNode, value: T, seen: IdHashSet, term: Term, tail: ConsList): PathNode { + val deref = if (seen.contains(deref(term))) term else deref(term) val arguments = deref.arguments() - val newTail = arguments + tail + val newTail = arguments.reversed().fold(tail) { list, t -> list.prepend(t) } val nextNode = node.nextOrDefault(symbolOrWildcard(deref)) { sym -> PathNode(sym, arguments.size) } @@ -84,21 +81,21 @@ class TermTrie() { node.putNext(nextNode.addValue(value)) } else { - node.putNext(putValue(nextNode, value, newTail.first(), newTail.drop(1))) + node.putNext(putValue(nextNode, value, seen.add(term), newTail.first()!!, newTail.drop(1))) } } - private fun removeValue(node: PathNode, value: T, term: Term, tail: List): PathNode { - val deref = deref(term) + private fun removeValue(node: PathNode, value: T, seen: IdHashSet, term: Term, tail: ConsList): PathNode { + val deref = if (seen.contains(deref(term))) term else deref(term) val arguments = deref.arguments() - val newTail = arguments + tail + val newTail = arguments.reversed().fold(tail) { list, t -> list.prepend(t) } return node.next(symbolOrWildcard(deref))?.let { nextNode -> //invariant: terms arity is fixed assert(nextNode.arity == arguments.size) - return if (newTail.isEmpty()) { + return if (newTail.isEmpty) { val newNext = nextNode.removeValue(value) if (newNext != nextNode) { if (!newNext.hasValues()) { @@ -113,7 +110,7 @@ class TermTrie() { } } else { - val newNext = removeValue(nextNode, value, newTail.first(), newTail.drop(1)) + val newNext = removeValue(nextNode, value, seen.add(term), newTail.first()!!, newTail.drop(1)) if (newNext !== nextNode) { if (!newNext.hasNext()) { node.removeNext(nextNode) @@ -130,12 +127,12 @@ class TermTrie() { } ?: node } - private fun visitMatching(node: PathNode, term: Term, tail: List, visitor: (T) -> Unit) { - val deref = deref(term) + private fun visitMatching(node: PathNode, seen: IdHashSet, term: Term, tail: ConsList, visitor: (T) -> Unit) { + val deref = if (seen.contains(deref(term))) term else deref(term) val sym = symbolOrWildcard(deref) if (sym == WILDCARD) { - if (tail.size > 0) { - node.skipAllNext().forEach { visitMatching(it, tail.first(), tail.drop(1), visitor) } + if (!tail.isEmpty) { + node.skipAllNext().forEach { visitMatching(it, seen.add(term), tail.first()!!, tail.drop(1), visitor) } } else { node.allNext().forEach { visitAll(it, visitor) } @@ -144,15 +141,15 @@ class TermTrie() { } else { node.next(sym)?.let { nn -> nn.values().forEach(visitor) - val newTail = deref.arguments() + tail - if (newTail.isNotEmpty()) { - visitMatching(nn, newTail.first(), newTail.drop(1), visitor) + val newTail = deref.arguments().reversed().fold(tail) { list, t -> list.prepend(t) } + if (!newTail.isEmpty) { + visitMatching(nn, seen.add(term), newTail.first()!!, newTail.drop(1), visitor) } } node.next(WILDCARD)?.let { nn -> nn.values().forEach(visitor) - if (tail.isNotEmpty()) { - visitMatching(nn, tail.first(), tail.drop(1), visitor) + if (!tail.isEmpty) { + visitMatching(nn, seen.add(term), tail.first()!!, tail.drop(1), visitor) } } } diff --git a/reactor/Test/test/TestTermTrie.kt b/reactor/Test/test/TestTermTrie.kt index d3fc051e..550711f3 100644 --- a/reactor/Test/test/TestTermTrie.kt +++ b/reactor/Test/test/TestTermTrie.kt @@ -21,12 +21,12 @@ class TestTermTrie { val t3 = parse("f{g h{i k{l m{o p} n}}}") val t4 = parse("f{g h{i q }}") - val trie1 = TermTrie().run { - put(t1, "foo").run { - put(t2, "bar").run { - put(t3, "qux").run { - put(t4, "blah") - } } } } + val trie1 = TermTrie().runs( + { put(t1, "foo") }, + { put(t2, "bar") }, + { put(t3, "qux") }, + { put(t4, "blah") } + ) assertEquals(setOf("foo", "bar", "qux", "blah"), trie1.allValues().toSet()) assertEquals(setOf("bar"), trie1.lookupValues(t2).toSet()) @@ -52,12 +52,12 @@ class TestTermTrie { val t3 = parse("f{g h{i k{l m{o p} n}}}") val t4 = parse("f{g h{i q }}") - val trie1 = TermTrie().run { - put(t1, "foo").run { - put(t2, "bar").run { - put(t3, "qux").run { - put(t4, "blah") - } } } } + val trie1 = TermTrie().runs( + { put(t1, "foo") }, + { put(t2, "bar") }, + { put(t3, "qux") }, + { put(t4, "blah") } + ) assertEquals(setOf("foo", "bar", "qux", "blah"), trie1.allValues().toSet()) assertEquals(setOf("bar"), trie1.lookupValues(t2).toSet()) @@ -83,10 +83,10 @@ class TestTermTrie { val t1 = parse("a{b c}") val t2 = parse("a{b d}") - val trie1 = TermTrie().run { - put(t1, "foo").run { - put(t2, "bar") - } } + val trie1 = TermTrie().runs( + { put(t1, "foo") }, + { put(t2, "bar") } + ) assertEquals(setOf("foo", "bar"), trie1.allValues().toSet()) assertEquals(setOf("foo"), trie1.lookupValues(t1).toSet()) @@ -118,13 +118,13 @@ class TestTermTrie { val t4 = parse("f{g{X b} X}") val t5 = parse("f{X Y}") - val tt = TermTrie().run { - put(t1, "t1").run { - put(t2, "t2").run { - put(t3, "t3").run { - put(t4, "t4").run { - put(t5, "t5") - } } } } } + val tt = TermTrie().runs( + { put(t1, "t1") }, + { put(t2, "t2") }, + { put(t3, "t3") }, + { put(t4, "t4") }, + { put(t5, "t5") } + ) assertEquals(setOf("t3", "t5"), tt.lookupValues(parse("f{g{b c} a}")).toSet()) assertEquals(setOf("t3", "t4", "t5"), tt.lookupValues(parse("f{g{b X} a}")).toSet()) @@ -142,10 +142,10 @@ class TestTermTrie { val t1 = parse("a{b c}") val t2 = parse("b{c d}") - val trie1 = TermTrie().run { - put(t1, "foo").run { - put(t2, "bar") - } } + val trie1 = TermTrie().runs ( + { put(t1, "foo") }, + { put(t2, "bar") } + ) assertEquals(setOf("foo"), trie1.lookupValues(parse("a{b c d}")).toSet()) assertEquals(setOf("bar"), trie1.lookupValues(parse("b{c d e}")).toSet()) @@ -159,10 +159,10 @@ class TestTermTrie { val t1 = parse("a{X c}") val t2 = parse("b{c Y}") - val trie1 = TermTrie().run { - put(t1, "foo").run { - put(t2, "bar") - } } + val trie1 = TermTrie().runs( + { put(t1, "foo") }, + { put(t2, "bar") } + ) assertEquals(setOf("foo"), trie1.lookupValues(parse("a{b c}")).toSet()) assertEquals(setOf("foo"), trie1.lookupValues(parse("a{b c d}")).toSet()) @@ -185,12 +185,12 @@ class TestTermTrie { val t3 = parse("f{g h{i k{l m{o p} n}}}") val t4 = parse("f{g h{i q }}") - val trie1 = TermTrie().run { - put(t1, "foo").run { - put(t2, "bar").run { - put(t3, "qux").run { - put(t4, "blah") - } } } } + val trie1 = TermTrie().runs( + { put(t1, "foo") }, + { put(t2, "bar") }, + { put(t3, "qux") }, + { put(t4, "blah") } + ) assertEquals(setOf("foo"), trie1.lookupValues(parse("a{X c}")).toSet()) assertEquals(setOf("foo"), trie1.lookupValues(parse("a{b Y}")).toSet()) @@ -209,13 +209,13 @@ class TestTermTrie { val t4 = parse("f{g h{Z k{l m{o p} n}}}") val t5 = parse("f{Z h{i q }}") - val trie1 = TermTrie().run { - put(t1, "foo").run { - put(t2, "bar").run { - put(t3, "bazz").run { - put(t4, "qux").run { - put(t5, "blah") - } } } } } + val trie1 = TermTrie().runs( + { put(t1, "foo") }, + { put(t2, "bar") }, + { put(t3, "bazz") }, + { put(t4, "qux") }, + { put(t5, "blah") } + ) assertEquals(setOf("foo", "bar", "bazz"), trie1.lookupValues(parse("a{X c{d e}}")).toSet()) assertEquals(setOf("foo", "bar", "bazz"), trie1.lookupValues(parse("a{b Y}")).toSet()) @@ -238,12 +238,12 @@ class TestTermTrie { val t3 = parse("a{c Y}") val t4 = parse("a{X Y}") - val trie1 = TermTrie().run { - put(t1, "foo").run { - put(t2, "bar").run { - put(t3, "bazz").run { - put(t4, "qux") - } } } } + val trie1 = TermTrie().runs( + { put(t1, "foo") }, + { put(t2, "bar") }, + { put(t3, "bazz") }, + { put(t4, "qux") } + ) // value order/cardinality no longer maintained assertEquals(setOf("foo", "bar", "qux"), trie1.lookupValues(parse("a{b X}")).toSet()) @@ -259,22 +259,63 @@ class TestTermTrie { @Test fun testRefTerm() { - val list = parse("f {c nil}") val varRef = MockRef(MockVar("TAIL")) - val pattern = MockFun("g", MockFun("h"), MockFun("f", varRef, MockFun("nil"))) - - val trie1 = TermTrie().run { - put(parse("g {h f {a nil}}"), "bar").run { - put(parse("g {h f {a f {b nil}}}"), "bazz").run { - put(pattern, "foo").run { - put(parse("g {h foo {nil}}"), "qux") - } } } } - - val key = MockFun("g", MockFun("h"), list) + val trie1 = TermTrie().runs( + { put(parse("g {h f {a nil}}"), "bar") }, + { put(parse("g {h f {a f {b nil}}}"), "bazz") }, + { put(pattern, "foo") }, + { put(parse("g {h foo {nil}}"), "qux") } + ) assertEquals(listOf("qux"), trie1.lookupValues(parse("g {h foo {nil}}")).toList()) + + val key = MockFun("g", MockFun("h"), parse("f {c nil}")) assertEquals(listOf("foo"), trie1.lookupValues(key).toList()) + } + + @Test + fun testRefList() { + val trie = TermTrie().runs( + { put(parse("f{X Y}"), "foo") }, + { put(parse("f{a Z}"), "bar") }, + { put(parse("f{a f{b W}}"), "bazz") } + ) + + val f = parse("f {b nil}") + val key = MockFun("f", parse("a"), MockRef(MockRef(f))) + assertEquals(setOf("foo", "bar", "bazz"), trie.lookupValues(key).toSet()) + } + + + @Test + fun testCyclicTerm() { + val trie = TermTrie().runs( + { put(parse("cst{ @1 n{ a{b c} d{^1} } }"), "foo") }, + { put(parse("cst{ n{ @2 a{b ^2} d{e} } }"), "bar") } + ) + + assertEquals(setOf("foo", "bar"), trie.lookupValues(parse("cst{ N }")).toSet()) + assertEquals(setOf("foo", "bar"), trie.lookupValues(parse("cst{ n{ X Y } }")).toSet()) + assertEquals(setOf("foo", "bar"), trie.lookupValues(parse("cst{ n{ X d{Z} } }")).toSet()) + assertEquals(setOf("foo", "bar"), trie.lookupValues(parse("cst{ n{ a{b W} V } }")).toSet()) + assertEquals(listOf("bar"), trie.lookupValues(parse("cst{ n{ @1 a{b ^1} V } }"))) + assertEquals(listOf("foo"), trie.lookupValues(parse("cst{ @1 n{ a{b S} d{^1} } }"))) + + val trie2 = trie.runs( + { remove(parse("cst{ n{ @2 a{b ^2} d{e} } }"), "bar") } + ) + assertEquals(emptyList(), trie2.lookupValues(parse("cst{ n{ @1 a{b ^1} V } }"))) + assertEquals(listOf("foo"), trie2.lookupValues(parse("cst{ @1 n{ a{b S} d{^1} } }"))) } + + fun T.runs(vararg blocks: T.() -> T): T { + var t = this + for (blk in blocks) { + t = t.blk() + } + return t + } + } \ No newline at end of file