diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/core/MatchTrie.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/core/MatchTrie.kt index ff414964..6d8caf4d 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/MatchTrie.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/MatchTrie.kt @@ -236,7 +236,7 @@ internal class MatchTrie(val rule: Rule, when (arg) { is Logical<*> -> fromArgs.addAll(aux.forLogical(arg)) - is Term -> fromArgs.addAll(aux.forTerm(arg)) + is Term -> fromArgs.addAll(aux.forTermAndConstraint(arg, cst)) is Any -> fromArgs.addAll(aux.forValue(arg)) } diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/core/Matcher.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/core/Matcher.kt index 79aa52cf..ad7458d8 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/Matcher.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/Matcher.kt @@ -174,42 +174,3 @@ fun Any?.asTerm(): Term = when (this) { // FIXME: "unpacking" the logicals fun Term.toValue(): Any? = if (this is Constant) this.symbol() else this -abstract class TermImpl(val symbol: Any) : Term { - - override fun symbol() = symbol - - override fun compareTo(other: Term): Int { - return if (this.javaClass === other.javaClass) - symbol.toString().compareTo(symbol.toString()) - else - this.javaClass.toString().compareTo(other.javaClass.toString()) - } -} - -open class Function(symbol: Any, val arguments: List) : TermImpl(symbol) { - - override fun arguments(): Collection = arguments - - override fun `is`(kind: Term.Kind?): Boolean = (kind === Term.Kind.FUN) - - override fun get(): Term = this -} - -class Constant(symbol: Any) : Function(symbol, emptyList()) {} - -class WrapConstant(val orig: Term) : Function(orig, emptyList()) {} - -class WrapFreeLogical(val logical: Logical<*>) : Function(logical.findRoot(), emptyList()) {} - -class WrapGroundLogical(val logical: Logical) : - Function(logical.findRoot().value().symbol(), - logical.findRoot().value().arguments().toList()) {} - -class Variable(symbol: Any) : TermImpl(symbol) { - - override fun arguments(): Collection = emptyList() - - override fun `is`(kind: Term.Kind): Boolean = (kind === Term.Kind.VAR) - - override fun get(): Term = TODO() -} \ No newline at end of file diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/core/OccurrenceStore.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/core/OccurrenceStore.kt index 73ececba..b27a9d4a 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/OccurrenceStore.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/OccurrenceStore.kt @@ -46,6 +46,8 @@ interface OccurrenceIndex { fun forTerm(term: Term): Iterable + fun forTermAndConstraint(term: Term, cst: Constraint): Iterable + fun forValue(value: Any): Iterable } @@ -57,13 +59,13 @@ class OccurrenceStore : LogicalObserver, OccurrenceIndex { val currentFrame: () -> StoreHolder - lateinit var symbol2occurrences: PersMap> + var symbol2occurrences: PersMap> - lateinit var logical2occurrences: PersMap>, IdHashSet> + var logical2occurrences: PersMap>, IdHashSet> - lateinit var term2occurrences: TermTrie + var term2occurrences: TermTrie - lateinit var value2occurrences: PersMap> + var value2occurrences: PersMap> constructor(copyFrom: OccurrenceStore, currentFrame: () -> StoreHolder) { @@ -88,7 +90,7 @@ class OccurrenceStore : LogicalObserver, OccurrenceIndex { when (value) { is Term -> { for (occ in toMerge) { - this.term2occurrences = term2occurrences.put(value, occ) + this.term2occurrences = term2occurrences.put(value.withConstraint(occ.constraint()), occ) } } is Any -> { @@ -144,7 +146,7 @@ class OccurrenceStore : LogicalObserver, OccurrenceIndex { currentFrame().addObserver(value) { frame -> frame.store() } } is Term -> { - this.term2occurrences = term2occurrences.put(value, occ) + this.term2occurrences = term2occurrences.put(value.withConstraint(occ.constraint()), occ) } is Any -> { this.value2occurrences = value2occurrences.put(value, @@ -222,15 +224,37 @@ class OccurrenceStore : LogicalObserver, OccurrenceIndex { } override fun forTerm(term: Term): Iterable { - return term2occurrences.lookupValues(term).filter { it.isStored() } + return term2occurrences.lookupValues(term.withAny()).filter { it.isStored() } + } + + override fun forTermAndConstraint(term: Term, cst: Constraint): Iterable { + return term2occurrences.lookupValues(term.withConstraint(cst)).filter { it.isStored() } } override fun forValue(value: Any): Iterable { return (value2occurrences[value] ?: emptySet()).filter { co -> co.isStored() } } + private fun Term.withConstraint(cst: Constraint): Term = Function(CONSTRAINT, listOf(Function(cst.symbol(), emptyList()), this)) + + private fun Term.withAny(): Term = Function(CONSTRAINT, listOf(Variable(ANY), this)) + + private companion object { + + private val CONSTRAINT = object : Any() { + override fun toString(): String = "CONSTRAINT" + } + + private val ANY = object : Any() { + override fun toString(): String = "ANY" + } + + } + } + + private data class MemConstraintOccurrence(val currentFrame: () -> HandlerFrame, val constraint: Constraint, val arguments: List<*>) : ConstraintOccurrence, LogicalObserver, diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/core/TermImpl.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/core/TermImpl.kt new file mode 100644 index 00000000..55bb4ef7 --- /dev/null +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/TermImpl.kt @@ -0,0 +1,44 @@ +package jetbrains.mps.logic.reactor.core + +import jetbrains.mps.logic.reactor.logical.Logical +import jetbrains.mps.unification.Term + +abstract class TermImpl(val symbol: Any) : Term { + + override fun symbol() = symbol + + override fun compareTo(other: Term): Int { + return if (this.javaClass === other.javaClass) + symbol.toString().compareTo(symbol.toString()) + else + this.javaClass.toString().compareTo(other.javaClass.toString()) + } +} + +open class Function(symbol: Any, val arguments: List) : TermImpl(symbol) { + + override fun arguments(): Collection = arguments + + override fun `is`(kind: Term.Kind?): Boolean = (kind === Term.Kind.FUN) + + override fun get(): Term = this +} + +class Constant(symbol: Any) : Function(symbol, emptyList()) {} + +class WrapConstant(val orig: Term) : Function(orig, emptyList()) {} + +class WrapFreeLogical(val logical: Logical<*>) : Function(logical.findRoot(), emptyList()) {} + +class WrapGroundLogical(val logical: Logical) : + Function(logical.findRoot().value().symbol(), + logical.findRoot().value().arguments().toList()) {} + +class Variable(symbol: Any) : TermImpl(symbol) { + + override fun arguments(): Collection = emptyList() + + override fun `is`(kind: Term.Kind): Boolean = (kind === Term.Kind.VAR) + + override fun get(): Term = TODO() +} \ No newline at end of file diff --git a/reactor/Test/test/TestMatcher.kt b/reactor/Test/test/TestMatcher.kt index 397d9f0a..7d83d5b5 100644 --- a/reactor/Test/test/TestMatcher.kt +++ b/reactor/Test/test/TestMatcher.kt @@ -1,6 +1,7 @@ import jetbrains.mps.logic.reactor.core.* import jetbrains.mps.logic.reactor.evaluation.ConstraintOccurrence import jetbrains.mps.logic.reactor.logical.Logical +import jetbrains.mps.logic.reactor.program.Constraint import jetbrains.mps.logic.reactor.program.ConstraintSymbol import jetbrains.mps.unification.Term import jetbrains.mps.unification.Unification @@ -32,6 +33,11 @@ class TestMatcher { co.arguments().any { it is Term && Unification.unify(it, term).isSuccessful } } + override fun forTermAndConstraint(term: Term, cst: Constraint): Iterable = + stored.filter { co -> + co.constraint().symbol() == cst.symbol() && co.arguments().any { it is Term && Unification.unify(it, term).isSuccessful } + } + override fun forValue(value: Any): Iterable = stored.filter { co -> co.arguments().contains(value) } }