diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/core/Handler.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/core/Handler.kt index 782f4299..29e3d093 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/Handler.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/Handler.kt @@ -40,7 +40,7 @@ class Handler { stored.filter { co -> constraint.matches(co) && acceptable(co) } } - val match = matcher.lookupMatches(active).find { pm -> pm.rule.guard().all { prd -> askPredicate(prd) } } + val match = matcher.lookupMatches(active).find { pm -> pm.rule.checkGuard(pm.logicalContext()) } if (match != null) { for ((cst, occ) in match.discarded) { @@ -68,17 +68,14 @@ class Handler { } } - private fun askPredicate(predicate: Predicate): Boolean = - sessionSolver.ask(predicate.symbol(), * predicate.arguments().toTypedArray()) + private fun Rule.checkGuard(logicalContext: LogicalContext): Boolean = + guard().all { prd -> askPredicate(prd.invocation(logicalContext)) } + + private fun askPredicate(invocation: PredicateInvocation): Boolean = + sessionSolver.ask(invocation.predicate().symbol(), * invocation.arguments().toTypedArray()) private fun tellPredicate(invocation: PredicateInvocation) { sessionSolver.tell(invocation.predicate().symbol(), * invocation.arguments().toTypedArray()) } } - -interface HandlingContext { - - fun substitute(logicalPattern: LogicalPattern<*>) : Any - -} \ No newline at end of file diff --git a/reactor/Test/test/LogicalHelper.kt b/reactor/Test/test/LogicalHelper.kt index c2029468..215d6e2e 100644 --- a/reactor/Test/test/LogicalHelper.kt +++ b/reactor/Test/test/LogicalHelper.kt @@ -30,8 +30,11 @@ inline fun logicalPattern(name1: String, name2: String, name3: fun Logical.get(): T = findRoot().value() -fun TestLogical.set(t: T) { - find().value = t +fun Logical.set(t: T) { + if (this is TestLogical) + find().value = t + else + throw IllegalStateException("unexpected receiver $this") } data class TestLogical(val name: String, var value: T?, var parent: TestLogical?) : Logical { @@ -69,6 +72,8 @@ data class TestLogical(val name: String, var value: T?, var parent: TestLogic fun union(other: TestLogical) { if (find() != other.find()) find().parent = other } + + override fun toString(): String = "$name(^${parent?.name ?: null})=$value" } data class TestLogicalPattern(val name: String, val type: Class) : LogicalPattern { diff --git a/reactor/Test/test/RulesHelper.kt b/reactor/Test/test/RulesHelper.kt index 63178c07..eb0679e2 100644 --- a/reactor/Test/test/RulesHelper.kt +++ b/reactor/Test/test/RulesHelper.kt @@ -150,6 +150,6 @@ private data class TestConstraintOccurrence(val constraint: Constraint, val argu override fun arguments(): Collection = arguments - override fun toString(): String = "#${constraint().symbol()}(${arguments().joinToString()})#${id}" + override fun toString(): String = "${constraint().symbol()}(${arguments().joinToString()})" } diff --git a/reactor/Test/test/TestProgram.kt b/reactor/Test/test/TestProgram.kt index 3b685df6..16a49ffb 100644 --- a/reactor/Test/test/TestProgram.kt +++ b/reactor/Test/test/TestProgram.kt @@ -80,10 +80,47 @@ class TestProgram { ) ).session("logicalValue").run { assertEquals(setOf(ConstraintSymbol("foo", 1), ConstraintSymbol("bar", 1)), constraintSymbols()) + assertEquals(2, constraintOccurrences().count()) val yval = constraintOccurrences(ConstraintSymbol("bar", 1)).first().arguments().first() assertEquals(66, (yval as Logical).get()) } } + @Test + fun simpleProgram() { + val (X, Y) = logicalPattern("X", "Y") + + program( + rule("main", + headReplaced( + constraint("main") + ), + body( + statement({ x -> x.set(5) }, X), + constraint("val", X) + ) + ), + rule("dec", + headReplaced( + constraint("val", X) + ), + guard( + expression({ x -> x.get() > 0 }, X) + ), + body( + constraint("trail", X), + statement({ x, y -> y.set(x.get() - 1)}, X, Y), + constraint("val", Y) + ) + ) + ).session("dec").run { + assertEquals(setOf(ConstraintSymbol("val", 1), ConstraintSymbol("trail", 1)), constraintSymbols()) + assertEquals(1, constraintOccurrences(ConstraintSymbol.symbol("val", 1)).count()) + val a = constraintOccurrences(ConstraintSymbol.symbol("val", 1)).first().arguments().first() + assertEquals(0, (a as Logical).get()) + assertEquals(5, constraintOccurrences(ConstraintSymbol.symbol("trail", 1)).count()) + println(constraintOccurrences(ConstraintSymbol.symbol("trail", 1))) + } + } }