diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleIndex.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleIndex.kt index f21627c0..098f9ef0 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleIndex.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleIndex.kt @@ -246,7 +246,7 @@ class RuleIndex(): Iterable, RuleLookup // all values should be accepted by a meta logical wildcardSelectors[argIdx].set(ruleBit) is Term -> - termSelectors.set(argIdx, termSelectors[argIdx].put(arg, ruleBit to headPos)) + termSelectors[argIdx].put(arg, ruleBit to headPos) is Any -> value2indices.getOrPut(arg) { hashSetOf() }.add(ruleBit to headPos) else -> @@ -264,7 +264,7 @@ class RuleIndex(): Iterable, RuleLookup // all values should be accepted by a meta logical wildcardSelectors[argIdx].clear(ruleBit) is Term -> - termSelectors.set(argIdx, termSelectors[argIdx].remove(arg, ruleBit to headPos)) + termSelectors[argIdx].remove(arg, ruleBit to headPos) is Any -> value2indices.get(arg)?.remove(ruleBit to headPos) else -> diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/util/ClassicIndexedTermTrie.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/util/ClassicIndexedTermTrie.kt new file mode 100644 index 00000000..17e033f8 --- /dev/null +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/util/ClassicIndexedTermTrie.kt @@ -0,0 +1,435 @@ +/* + * Copyright 2014-2018 JetBrains s.r.o. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package jetbrains.mps.logic.reactor.util + +import gnu.trove.map.hash.TIntIntHashMap +import gnu.trove.map.hash.TIntObjectHashMap +import jetbrains.mps.logic.reactor.logical.VarSymbol +import jetbrains.mps.unification.Term +import java.util.* +import kotlin.collections.ArrayList + +/** + * @author Fedor Isakov + */ + +/** + * A map-like structure to keep multiple values associated with a single term. + * + * Term variables are supported: a variable matches any term, including other variable. All variables are treated + * as "wildcards", so the subsequent unification of the query term and the key does not necessarily succeeds. + * + * To access the stored values, the method `lookupValues` returns all values associated with terms. + * The method `allValues` returns all stored values. Neither method makes any guarantees about the cardinality or + * the order of the values returned. + * + * The implementation is based on the structure called "discrimination tree" as described in the chapter 26 of [1]. + * + * The indexed variant of the term trie adds in addiiton a possibility to index all values by an int value. + * An index may be specified in all put/remove and also to lookup methods, and limits the scope of + * internal nodes to be traversed, thus eliminating unnecessary calls. + * + * [1] Alan Robinson and Andrei Voronkov (Eds.). 2001. Handbook of Automated Reasoning. + * Elsevier Sci. Pub. B. V., Amsterdam, The Netherlands, The Netherlands. + * An implementation of TermTrie not relying on persistent data structures. + */ + +class ClassicIndexedTermTrie : IndexedTermTrie { + + private companion object { + + private val WILDCARD = object : Any() { + override fun toString() = "WILDCARD" + } + + private val KEYHOLDER = object : Any() { + override fun toString() = "KEYHOLDER" + } + + } + + private var root: PathNode + + init { + root = PathNode(WILDCARD, 1) + } + + override fun put(term: Term, value: T) { + putValue(term, -1, value) + } + + override fun put(term: Term, index: Int, value: T) { + putValue(term, index, value) + } + + override fun remove(term: Term, value: T) { + removeValue(term, -1, value) + } + + override fun remove(term: Term, index: Int, value: T) { + removeValue(term, index, value) + } + + override fun lookupValues(term: Term): Iterable { + val result = ArrayList() + visitMatching(term, root.allIndexMask()) { value, index -> result.add(value) } + return result + } + + override fun lookupValues(term: Term, indexMask: IndexMask): Iterable { + val result = ArrayList() + visitMatching(term, indexMask) { value, index -> result.add(value) } + return result + } + + override fun forValuesWithIndex(term: Term, indexMask: IndexMask?, callback: (T, Int) -> Unit) { + visitMatching(term, indexMask) { value, index -> callback(value, index) } + } + + override fun allValues(): Iterable { + val result = ArrayList() + visitAll(root, root.allIndexMask(), { value, index -> result.add(value) }) + return result + } + + override fun allValues(indexMask: IndexMask): Iterable { + val result = ArrayList() + visitAll(root, indexMask, { value, index -> result.add(value) }) + return result + } + + private fun putValue(matchTerm: Term, index: Int, value: T) { + val seen = IdentityHashMap() + val nodeStack = arrayListOf(root) + val termStack = arrayListOf(matchTerm) + val argPool = arrayListOf() + while (!termStack.isEmpty()) { + val node = nodeStack.peek() + val term = termStack.pop() + + node.addIndex(index) + + // dereferece the term only if it hasn't been dereferenced before + val deref = deref(term).let { dt -> seen[dt]?.run { term } ?: dt.apply { seen[dt] = term } } + argPool.addAll(deref.arguments()) + for(i in 1..argPool.size) { + termStack.push(argPool[argPool.size - i]) + } + val nextNode = node.nextOrDefault(symbolOrWildcard(deref)) { sym -> PathNode(sym, argPool.size) } + nodeStack.push(nextNode) + argPool.clear() + } + + val head = nodeStack.pop() + head.addValue(value, index) + } + + private fun removeValue(matchTerm: Term, index: Int, value: T) { + val seen = IdentityHashMap() + val nodeStack = arrayListOf(root) + val termStack = arrayListOf(matchTerm) + val argPool = arrayListOf() + while (!termStack.isEmpty()) { + val node = nodeStack.peek() + val term = termStack.pop() + + // dereferece the term only if it hasn't been dereferenced before + val deref = deref(term).let { dt -> seen[dt]?.run { term } ?: dt.apply { seen[dt] = term } } + argPool.addAll(deref.arguments()) + for(i in 1..argPool.size) { + termStack.push(argPool[argPool.size - i]) + } + argPool.clear() + + val nextNode = node.next(symbolOrWildcard(deref)) + if (nextNode == null) { + // term not found + return + } + nodeStack.push(nextNode) + } + + var head = nodeStack.pop() + if (head.removeValue(value, index)) { + while (nodeStack.isNotEmpty()) { + val base = nodeStack.pop() + base.removeIndex(index) + if (head.isLeaf()) { + base.dropNext(head) + } + head = base + } + } + } + + /** + * Call the visitor function with values of the given node and all its direct and indirect successors. + */ + private fun visitAll(node: PathNode, indexMask: IndexMask, visitor: (T, Int) -> Unit) { + node.forEachValueWithIndex(indexMask, visitor) + node.allNext(indexMask).forEach { visitAll(it, indexMask, visitor) } + } + + private fun collectAllLeaves(node: PathNode, indexMask: IndexMask?, leafs: MutableList>) { + val nodeStack = arrayListOf(node) + while (nodeStack.isNotEmpty()) { + val n = nodeStack.pop() + if (n.isLeaf()) leafs.add(n) + n.allNext(indexMask).forEach { + nodeStack.push(it) + } + } + } + + private fun deref(term: Term): Term { + var deref = term + while (deref.`is`(Term.Kind.REF)) { + deref = deref.get() + } + return deref + } + + private fun symbolOrWildcard(term: Term): Any { + return if (term.`is`(Term.Kind.VAR) || (term.`is`(Term.Kind.FUN) && term.symbol() is VarSymbol)) { + WILDCARD + + } else { + term.symbol() ?: throw NullPointerException("term symbol can't be null") + } + } + + /** + * Given a pattern term, which may contain variables that are treated as wildcards, + * call the passed visitor function with values of all matching nodes and all the nodes that precede them. + */ + private fun visitMatching(pattern: Term, indexMask: IndexMask?, visitor: (T, Int) -> Unit) { + val seenNonLeaf = IdentityHashMap() + val canonicTerms = IdentityHashMap>() + val visitStack = arrayListOf(listOf(root) to listOf(pattern)) + val allLeaves = arrayListOf>() + + while (!visitStack.isEmpty()) { + val (bases, ptnTerms) = visitStack.pop() + if (!ptnTerms.isEmpty()) { + val ptnHead = ptnTerms.first() + val ptnTail = ptnTerms.subList(1, ptnTerms.size) + + if (!canonicTerms.containsKey(ptnHead)) { + // dereferece the term only if it hasn't been dereferenced before + val derefPtnHead = deref(ptnHead).let { dt -> seenNonLeaf[dt]?.run { ptnHead } ?: dt } + canonicTerms[ptnHead] = derefPtnHead to symbolOrWildcard(derefPtnHead) + } + val (term, sym) = canonicTerms[ptnHead] !! + + for (base in bases) { + if (sym == WILDCARD) { + if (!ptnTail.isEmpty()) { + // skip the current node + // match the patterns tail with the current node's direct successors + val (wcdTerms, wcdBases) = base.terms2bases(indexMask) + wcdTerms.forEach { if (it.isLeaf()) allLeaves.add(it) /*it.forEachValueWithIndex(indexMask, visitor)*/ } + visitStack.push(wcdBases to ptnTail) + + } else { + // wildcard consumes the rest of the trie + base.allNext(indexMask).forEach { collectAllLeaves(it, indexMask, allLeaves) } + } + + } else { + base.next(WILDCARD)?.let { nn -> + if (nn.isLeaf()) allLeaves.add(nn) +// nn.forEachValueWithIndex(indexMask, visitor) + if (!ptnTail.isEmpty()) { + visitStack.push(listOf(nn) to ptnTail) + } + } + base.next(sym)?.let { nn -> + if (nn.isLeaf()) allLeaves.add(nn) +// nn.forEachValueWithIndex(indexMask, visitor) + + // prepend this patternTerm's arguments to the tail pattern terms + val newTail = ArrayList(term.arguments()) + if (newTail.size > 0) { + seenNonLeaf[term] = term + } + newTail.addAll(ptnTail) + if (!newTail.isEmpty()) { + visitStack.push(listOf(nn) to newTail) + } + } + } + } + } + } + + indexMask?.forEach { index -> + allLeaves.forEach { it.forEachValueWithIndex(index, visitor) } + true + } ?: run { + allLeaves.forEach { it.forEachValue(visitor) } + } + } + + /** + * A trie node. Corresponds to a particular subterm. + */ + private class PathNode(val symbol: Any, + val arity: Int) + { + private val indexCardinalities = TIntIntHashMap() // map of index cardinalities + + private val next = HashMap>(8) + + private val indexedValues = TIntObjectHashMap>() + + fun allIndexMask(): IndexMask = indexCardinalities.keySet() + + fun addIndex(index: Int) { + val card = if (indexCardinalities.contains(index)) indexCardinalities.get(index) else 0 + indexCardinalities.put(index, card + 1) + } + + fun removeIndex(index: Int) { + val card = if (indexCardinalities.contains(index)) indexCardinalities.get(index) else 0 + assert(card > 0) + if (card > 1) indexCardinalities.put(index, card - 1) else indexCardinalities.remove(index) + } + + fun forEachValue(callback: (T, Int) -> Unit ) { + indexedValues.forEachEntry { index, values -> + values.forEach { callback(it, index) } + true + } + } + + fun forEachValueWithIndex(indexMask: IndexMask, callback: (T, Int) -> Unit ) { + indexMask.forEach { index -> + if (indexedValues.containsKey(index)) { indexedValues.get(index).forEach { callback(it, index) } } + true + } + } + + fun forEachValueWithIndex(index: Int, callback: (T, Int) -> Unit ) { + if (indexedValues.containsKey(index)) { indexedValues.get(index).forEach { callback(it, index) } } + } + + fun isLeaf(): Boolean = next.isEmpty() + + fun next(symbol: Any): PathNode? = next[symbol] + + /** + * Returns all trie nodes that are direct successors of this one. + */ + fun allNext(indexMask: IndexMask?): List> { + val res = arrayListOf>() + next.values.forEach { n -> + // containsAny ~~ NOT( all( NOT(contains)) + if (indexMask != null) { + if (!n.indexCardinalities.forEach { idx -> !indexMask.contains(idx) }) { + res.add(n) + } + } else { + res.add(n) + } + } + return res + } + + /** + * Helper method for processing a wildcard in pattern. + * + * Returns a pair of iterables: + * - first component contains all nodes that make up the next *term*; + * - second component contains base nodes that precede terms following that *term*. + * + * The trie keeps the terms _flattened_, and the following proposition holds. + * + * Let t = f(...) be a term. Let flt: Term->List be a function that transforms terms to lists of symbols. + * Let ar: Symbol->int be a function that returns arity for a given symbol. + * + * Then size(flt t) = 1 + sum . (map ar) (flt t), where '.' stands for function composition. + * + * The size of a list representing a flattened term is equal to + * the sum of arities of all symbols in this list plus 1. + */ + fun terms2bases(indexMask: IndexMask?): Pair>, List>> { + + // for every current node there is a number + // initially 0 + // counting down with every call to allNext() + // increased by current node's arity + + val terms = ArrayList>() + val bases = ArrayList>() + + val stack = arrayListOf(allNext(indexMask) to 0) + + while (stack.isNotEmpty()) { + val (nn, count) = stack.pop() + for (n in nn) { + val newCount = count + n.arity + if (newCount == 0) { + bases.add(n) + + } else { + terms.add(n) + stack.push(n.allNext(indexMask) to (newCount - 1)) + } + } + } + + return terms to bases + } + + + fun addValue(value: T, index: Int) { + addIndex(index) + (indexedValues.get(index) ?: hashSetOf().also { indexedValues.put(index, it) }).add(value) + } + + fun removeValue(value: T, index: Int): Boolean { + if (indexedValues.get(index)?.remove(value) ?: false) { + if (indexedValues.get(index)?.isEmpty() ?: false) { + indexedValues.remove(index) + } + removeIndex(index) + return true + } + return false + } + + inline fun nextOrDefault(symbol: Any, + default: (sym: Any) -> PathNode): PathNode + { + val existing = next[symbol] + if (existing != null) { + return existing + } else { + val default = default(symbol) + next[symbol] = default + return default + } + } + + fun dropNext(drop: PathNode) { + this.next.remove(drop) + } + +// override fun toString(): String = +// "${symbol}/${arity} (${values.keys.joinToString(", ")}) > [${next.map { it.symbol }.joinToString(", ")}] " + } +} diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/util/ClassicTermTrie.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/util/ClassicTermTrie.kt index 56cf1d21..ebe93a99 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/util/ClassicTermTrie.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/util/ClassicTermTrie.kt @@ -16,6 +16,7 @@ package jetbrains.mps.logic.reactor.util +import gnu.trove.map.hash.TIntIntHashMap import jetbrains.mps.logic.reactor.logical.VarSymbol import jetbrains.mps.unification.Term import java.util.* @@ -62,14 +63,12 @@ class ClassicTermTrie : TermTrie { root = PathNode(WILDCARD, 1) } - override fun put(term: Term, value: T): TermTrie { + override fun put(term: Term, value: T) { putValue(term, value) - return this } - override fun remove(term: Term, value: T): TermTrie { + override fun remove(term: Term, value: T) { removeValue(term, value) - return this } override fun lookupValues(term: Term): Iterable { @@ -235,6 +234,8 @@ class ClassicTermTrie : TermTrie { private class PathNode(val symbol: Any, val arity: Int) { + private val index = TIntIntHashMap() // map of index cardinalities + private val next = HashMap>(8) private val values = IdentityHashMap() diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/util/IndexedTermTrie.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/util/IndexedTermTrie.kt new file mode 100644 index 00000000..94341f5b --- /dev/null +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/util/IndexedTermTrie.kt @@ -0,0 +1,44 @@ +/* + * Copyright 2014-2021 JetBrains s.r.o. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package jetbrains.mps.logic.reactor.util + +import gnu.trove.set.TIntSet +import gnu.trove.set.hash.TIntHashSet +import jetbrains.mps.unification.Term + +/** + * @author Fedor Isakov + */ + +typealias IndexMask = TIntSet +fun indexMaskOf(vararg indices: Int): TIntSet = TIntHashSet(intArrayOf(*indices)) + +interface IndexedTermTrie : TermTrie { + + fun put(term: Term, index: Int, value: T) + + fun remove(term: Term, index: Int, value: T) + + fun lookupValues(term: Term, indexMask: IndexMask): Iterable + + fun forValuesWithIndex(term: Term, indexMask: IndexMask?, callback: (T, Int) -> Unit) + + fun allValues(indexMask: IndexMask): Iterable + +} + +fun indexedTermTrie(): IndexedTermTrie = ClassicIndexedTermTrie() diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/util/TermTrie.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/util/TermTrie.kt index a2f0d8eb..3b04abf0 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/util/TermTrie.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/util/TermTrie.kt @@ -38,13 +38,13 @@ import jetbrains.mps.unification.Term */ interface TermTrie { - fun put(term: Term, value: T): TermTrie + fun put(term: Term, value: T) - fun remove(term: Term, value: T): TermTrie + fun remove(term: Term, value: T) fun lookupValues(term: Term): Iterable fun allValues(): Iterable } -fun termTrie(): TermTrie = ClassicTermTrie() +fun termTrie(): TermTrie = ClassicIndexedTermTrie() diff --git a/reactor/Test/test/TestIndexedTermTrie.kt b/reactor/Test/test/TestIndexedTermTrie.kt new file mode 100644 index 00000000..a8526d4b --- /dev/null +++ b/reactor/Test/test/TestIndexedTermTrie.kt @@ -0,0 +1,144 @@ +import jetbrains.mps.logic.reactor.util.indexMaskOf +import jetbrains.mps.logic.reactor.util.indexedTermTrie +import jetbrains.mps.logic.reactor.util.termTrie +import jetbrains.mps.unification.test.MockTermsParser +import jetbrains.mps.unification.test.MockTermsParser.* +import org.junit.Assert +import org.junit.Assert.* +import org.junit.Test + +/* + * Copyright 2014-2021 JetBrains s.r.o. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +/** + * @author Fedor Isakov + */ +class TestIndexedTermTrie { + + @Test + fun testPut() { + val t1 = parseTerm("a{b c}") + val t2 = parseTerm("a{d e}") + val t3 = parseTerm("f{g h{i k{l m{o p} n}}}") + val t4 = parseTerm("f{g h{i q }}") + + val trie1 = indexedTermTrie().runs( + { put(t1, 1, "foo") }, + { put(t2, 2, "bar") }, + { put(t3, 2, "qux") }, + { put(t4, 3,"blah") } + ) + + assertEquals(setOf("foo", "bar", "qux", "blah"), trie1.allValues().toSet()) + assertEquals(setOf("bar", "qux"), trie1.allValues(indexMaskOf(2)).toSet()) + assertEquals(setOf("bar"), trie1.lookupValues(t2).toSet()) + assertEquals(setOf("bar"), trie1.lookupValues(t2, indexMaskOf(2)).toSet()) + assertEquals(setOf(), trie1.lookupValues(t2, indexMaskOf(3)).toSet()) + + val trie2 = trie1.also { it.put(t2, 3, "bazz") } + assertEquals(setOf("foo", "bar", "qux", "blah", "bazz"), trie2.allValues().toSet()) + assertEquals(setOf("blah", "bazz"), trie2.allValues(indexMaskOf(3)).toSet()) + assertEquals(setOf("foo", "blah", "bazz"), trie2.allValues(indexMaskOf(1, 3)).toSet()) + + assertEquals(setOf("bar", "bazz"), trie2.lookupValues(t2).toSet()) + assertEquals(setOf("bar", "bazz"), trie2.lookupValues(t2, indexMaskOf(2,3)).toSet()) + assertEquals(setOf(), trie2.lookupValues(t2, indexMaskOf(1)).toSet()) + assertEquals(setOf(), trie2.lookupValues(parseTerm("a{d f}")).toSet()) + assertEquals(setOf("qux"), trie2.lookupValues(t3).toSet()) + assertEquals(setOf("qux"), trie2.lookupValues(t3, indexMaskOf(2)).toSet()) + assertEquals(setOf("blah"), trie2.lookupValues(t4).toSet()) + assertEquals(setOf("blah"), trie2.lookupValues(t4, indexMaskOf(3)).toSet()) + + val trie3 = trie2.also { it.put(t4, 1, "shmoo") } + assertEquals(setOf("foo", "blah", "bazz", "shmoo"), trie2.allValues(indexMaskOf(1, 3)).toSet()) + assertEquals(setOf(), trie2.lookupValues(t2, indexMaskOf(1)).toSet()) + assertEquals(setOf("foo", "bar", "qux", "blah", "bazz", "shmoo"), trie3.allValues().toSet()) + assertEquals(setOf("qux", "blah", "bazz", "bar"), trie3.allValues(indexMaskOf(2, 3)).toSet()) + assertEquals(setOf("blah", "shmoo"), trie3.lookupValues(t4, indexMaskOf(1, 3)).toSet()) + assertEquals(setOf("shmoo"), trie3.lookupValues(t4, indexMaskOf(1)).toSet()) + } + + @Test + fun testRemove() { + val t1 = parseTerm("a{b c}") + val t2 = parseTerm("a{d e}") + val t3 = parseTerm("f{g h{i k{l m{o p} n}}}") + val t4 = parseTerm("f{g h{i q }}") + + val trie1 = indexedTermTrie().runs( + { put(t1, 1, "foo") }, + { put(t2, 1, "bar") }, + { put(t3, 2, "qux") }, + { put(t4, 2, "blah") } + ) + + assertEquals(setOf("foo", "bar", "qux", "blah"), trie1.allValues().toSet()) + assertEquals(setOf("bar"), trie1.lookupValues(t2).toSet()) + + val trie2 = trie1.also { it.remove(t2, "bazz") } + assertEquals(setOf("foo", "bar", "qux", "blah"), trie2.allValues().toSet()) + assertEquals(setOf("bar"), trie2.lookupValues(t2).toSet()) + + val trie3 = trie2.also { it.remove(t2, 1, "bar") } + assertEquals(setOf("foo", "qux", "blah"), trie3.allValues().toSet()) + assertEquals(setOf(), trie3.lookupValues(t2).toSet()) + + + val trie4 = trie3.also { it.remove(t3, 2, "qux") } + assertEquals(setOf("foo", "blah"), trie4.allValues().toSet()) + assertEquals(setOf(), trie4.lookupValues(t3).toSet()) + assertEquals(setOf("blah"), trie4.lookupValues(t4).toSet()) + } + + @Test + fun testRemovePut() { + val t1 = parseTerm("a{b c}") + val t2 = parseTerm("a{b d}") + + val trie1 = indexedTermTrie().runs( + { put(t1, 1, "foo") }, + { put(t2, 2, "bar") } + ) + + assertEquals(setOf("foo", "bar"), trie1.allValues().toSet()) + assertEquals(setOf("foo"), trie1.lookupValues(t1).toSet()) + assertEquals(setOf("bar"), trie1.lookupValues(t2).toSet()) + + val trie2 = trie1.also { it.remove(t2, 2, "bar") } + assertEquals(setOf("foo"), trie2.allValues().toSet()) + assertEquals(setOf("foo"), trie2.lookupValues(t1).toSet()) + assertEquals(setOf(), trie2.lookupValues(t2).toSet()) + + val trie3 = trie2.also { it.put(t2, 2, "bazz") } + assertEquals(setOf("foo", "bazz"), trie3.allValues().toSet()) + assertEquals(setOf("foo", "bazz"), trie3.allValues(indexMaskOf(1,2)).toSet()) + assertEquals(setOf("foo"), trie3.lookupValues(t1).toSet()) + assertEquals(setOf("bazz"), trie3.lookupValues(t2).toSet()) + + val trie4 = trie3.also { it.remove(t1, 1, "foo") } + assertEquals(setOf("bazz"), trie4.allValues().toSet()) + assertEquals(setOf(), trie4.lookupValues(t1).toSet()) + assertEquals(setOf("bazz"), trie4.lookupValues(t2).toSet()) + } + + fun T.runs(vararg blocks: T.() -> Unit): T { + for (blk in blocks) { + blk() + } + return this + } + +} \ No newline at end of file diff --git a/reactor/Test/test/TestTermTrie.kt b/reactor/Test/test/TestTermTrie.kt index ab1f4a09..f655c72e 100644 --- a/reactor/Test/test/TestTermTrie.kt +++ b/reactor/Test/test/TestTermTrie.kt @@ -28,7 +28,7 @@ class TestTermTrie { assertEquals(setOf("foo", "bar", "qux", "blah"), trie1.allValues().toSet()) assertEquals(setOf("bar"), trie1.lookupValues(t2).toSet()) - val trie2 = trie1.put(t2, "bazz") + val trie2 = trie1.also { it.put(t2, "bazz") } assertEquals(setOf("foo", "bar", "qux", "blah", "bazz"), trie2.allValues().toSet()) assertEquals(setOf("bar", "bazz"), trie2.lookupValues(t2).toSet()) @@ -36,7 +36,7 @@ class TestTermTrie { assertEquals(setOf("qux"), trie2.lookupValues(t3).toSet()) assertEquals(setOf("blah"), trie2.lookupValues(t4).toSet()) - val trie3 = trie2.put(t4, "shmoo") + val trie3 = trie2.also { it.put(t4, "shmoo") } assertEquals(setOf("foo", "bar", "qux", "blah", "bazz", "shmoo"), trie3.allValues().toSet()) assertEquals(setOf("blah", "shmoo"), trie3.lookupValues(t4).toSet()) } @@ -59,16 +59,16 @@ class TestTermTrie { assertEquals(setOf("foo", "bar", "qux", "blah"), trie1.allValues().toSet()) assertEquals(setOf("bar"), trie1.lookupValues(t2).toSet()) - val trie2 = trie1.remove(t2, "bazz") + val trie2 = trie1.also { it.remove(t2, "bazz") } assertEquals(setOf("foo", "bar", "qux", "blah"), trie2.allValues().toSet()) assertEquals(setOf("bar"), trie2.lookupValues(t2).toSet()) - val trie3 = trie2.remove(t2, "bar") + val trie3 = trie2.also { it.remove(t2, "bar") } assertEquals(setOf("foo", "qux", "blah"), trie3.allValues().toSet()) assertEquals(setOf(), trie3.lookupValues(t2).toSet()) - val trie4 = trie3.remove(t3, "qux") + val trie4 = trie3.also { it.remove(t3, "qux") } assertEquals(setOf("foo", "blah"), trie4.allValues().toSet()) assertEquals(setOf(), trie4.lookupValues(t3).toSet()) assertEquals(setOf("blah"), trie4.lookupValues(t4).toSet()) @@ -89,17 +89,17 @@ class TestTermTrie { assertEquals(setOf("foo"), trie1.lookupValues(t1).toSet()) assertEquals(setOf("bar"), trie1.lookupValues(t2).toSet()) - val trie2 = trie1.remove(t2, "bar") + val trie2 = trie1.also { it.remove(t2, "bar") } assertEquals(setOf("foo"), trie2.allValues().toSet()) assertEquals(setOf("foo"), trie2.lookupValues(t1).toSet()) assertEquals(setOf(), trie2.lookupValues(t2).toSet()) - val trie3 = trie2.put(t2, "bazz") + val trie3 = trie2.also { it.put(t2, "bazz") } assertEquals(setOf("foo", "bazz"), trie3.allValues().toSet()) assertEquals(setOf("foo"), trie3.lookupValues(t1).toSet()) assertEquals(setOf("bazz"), trie3.lookupValues(t2).toSet()) - val trie4 = trie3.remove(t1, "foo") + val trie4 = trie3.also { it.remove(t1, "foo") } assertEquals(setOf("bazz"), trie4.allValues().toSet()) assertEquals(setOf(), trie4.lookupValues(t1).toSet()) assertEquals(setOf("bazz"), trie4.lookupValues(t2).toSet()) @@ -127,7 +127,7 @@ class TestTermTrie { assertEquals(setOf("t3", "t4", "t5"), tt.lookupValues(parseTerm("f{g{b X} a}")).toSet()) assertEquals(setOf("t1", "t3", "t4", "t5"), tt.lookupValues(parseTerm("f{g{X Y} a}")).toSet()) - val tt2 = tt.remove(t4, "t4") + val tt2 = tt.also { it.remove(t4, "t4") } assertEquals(setOf("t3", "t5"), tt2.lookupValues(parseTerm("f{g{b X} a}")).toSet()) assertEquals(setOf("t1", "t3", "t5"), tt2.lookupValues(parseTerm("f{g{X Y} a}")).toSet()) @@ -247,8 +247,8 @@ class TestTermTrie { assertEquals(setOf("foo", "bar", "bazz", "qux"), trie1.lookupValues(parseTerm("a{X c}")).toSet()) assertEquals(setOf("foo", "bazz", "qux"), trie1.lookupValues(parseTerm("a{c c}")).toSet()) - val trie2 = trie1.remove(t2, "bar") - val trie3 = trie2.put(t2, "blah") + val trie2 = trie1.also { it.remove(t2, "bar") } + val trie3 = trie2.also { it.put(t2, "blah") } assertEquals(setOf("foo", "qux", "blah"), trie3.lookupValues(parseTerm("a{b X}")).toSet()) assertEquals(setOf("foo", "bazz", "qux", "blah"), trie3.lookupValues(parseTerm("a{X c}")).toSet()) @@ -320,12 +320,11 @@ class TestTermTrie { } - fun T.runs(vararg blocks: T.() -> T): T { - var t = this + fun T.runs(vararg blocks: T.() -> Unit): T { for (blk in blocks) { - t = t.blk() + blk() } - return t + return this } } \ No newline at end of file