Switched to using LogicalContext for the rule matching, tests

This commit is contained in:
Fedor Isakov 2015-12-22 11:08:00 +01:00
parent 8de8f09325
commit c04ecab6d0
5 changed files with 121 additions and 69 deletions

View File

@ -1,6 +1,8 @@
package jetbrains.mps.logic.reactor.core
import jetbrains.mps.logic.reactor.constraint.*
import jetbrains.mps.logic.reactor.logical.LogicalContext
import jetbrains.mps.logic.reactor.logical.LogicalPattern
import jetbrains.mps.logic.reactor.rule.Rule
import java.util.*
@ -11,19 +13,16 @@ import java.util.*
class Handler {
private val sessionSolver: SessionSolver
private val occurrenceFactory: (Constraint) -> ConstraintOccurrence
private val rules : MutableList<Rule> = ArrayList<Rule>()
private val stored : MutableList<ConstraintOccurrence> = ArrayList<ConstraintOccurrence>()
constructor(
sessionSolver: SessionSolver,
programRules: Iterable<Rule>,
occurrenceFactory: (Constraint) -> ConstraintOccurrence,
// for testing purposes only
occurrences: Iterable<ConstraintOccurrence>? = null)
{
this.sessionSolver = sessionSolver
this.occurrenceFactory = occurrenceFactory
this.rules.addAll(programRules)
if (occurrences != null) {
this.stored.addAll(occurrences)
@ -36,9 +35,9 @@ class Handler {
stored.add(active)
val matcher = object : Matcher(rules) {
override fun findOccurrences(constraint: Constraint, predicate: (ConstraintOccurrence) -> Boolean):
override fun findOccurrences(constraint: Constraint, acceptable: (ConstraintOccurrence) -> Boolean):
Iterable<ConstraintOccurrence> =
stored.filter { co -> constraint.matches(co) && predicate(co) }
stored.filter { co -> constraint.matches(co) && acceptable(co) }
}
val match = matcher.lookupMatches(active).find { pm -> pm.rule.guard().all { prd -> askPredicate(prd) } }
@ -47,8 +46,15 @@ class Handler {
for ((cst, occ) in match.discarded) {
discard(occ)
}
val lc = object: LogicalContext {
override fun valueFor(logicalPattern: LogicalPattern?): Any? {
throw UnsupportedOperationException()
}
}
for (item in match.rule.body()) {
activate(item)
activate(item, lc)
}
return false
@ -60,22 +66,24 @@ class Handler {
stored.remove(occ)
}
private fun activate(item: AndItem) {
private fun activate(item: AndItem, logicalContext: LogicalContext) {
when(item) {
is Constraint -> process(activate(item))
is Constraint -> process(item.occurrence(logicalContext))
is Predicate -> tellPredicate(item)
else -> throw IllegalArgumentException("unknown item ${item}")
}
}
private fun activate(constraint: Constraint): ConstraintOccurrence = occurrenceFactory(constraint)
private fun askPredicate(predicate: Predicate): Boolean =
sessionSolver.ask(predicate.symbol(), * predicate.arguments().toTypedArray())
private fun askPredicate(predicate: Predicate): Boolean {
return sessionSolver.ask(predicate.symbol(), * predicate.arguments().toTypedArray())
}
private fun tellPredicate(predicate: Predicate) {
private fun tellPredicate(predicate: Predicate) =
sessionSolver.tell(predicate.symbol(), * predicate.arguments().toTypedArray())
}
}
interface HandlingContext {
fun substitute(logicalPattern: LogicalPattern) : Any
}

View File

@ -4,11 +4,13 @@ import com.github.andrewoma.dexx.collection.ConsList
import jetbrains.mps.logic.reactor.constraint.Constraint
import jetbrains.mps.logic.reactor.constraint.ConstraintOccurrence
import jetbrains.mps.logic.reactor.logical.Logical
import jetbrains.mps.logic.reactor.logical.LogicalContext
import jetbrains.mps.logic.reactor.logical.LogicalPattern
import jetbrains.mps.logic.reactor.rule.Rule
import jetbrains.mps.unification.Term
import jetbrains.mps.unification.Unification
import java.lang.String
import java.util.*
/**
* @author Fedor Isakov
@ -43,7 +45,7 @@ abstract class Matcher(val rules: Collection<Rule>) {
return matchesFromKept + matchesFromDiscarded
}
abstract fun findOccurrences(constraint: Constraint, predicate: (ConstraintOccurrence) -> Boolean):
abstract fun findOccurrences(constraint: Constraint, acceptable: (ConstraintOccurrence) -> Boolean):
Iterable<ConstraintOccurrence>
}
@ -54,6 +56,7 @@ class PartialMatch(val rule: Rule) {
private set
var discarded = ConsList.empty<Pair<Constraint, ConstraintOccurrence>>()
private set
private lateinit var logicalContext : LogicalContext
private constructor(
original : PartialMatch,
@ -64,8 +67,6 @@ class PartialMatch(val rule: Rule) {
discarded = if (discard != null) original.discarded.append(discard) else original.discarded
}
fun clone() : PartialMatch = PartialMatch(this, null, null)
fun keep (constraint: Constraint, occ: ConstraintOccurrence) = PartialMatch(this, Pair(constraint, occ), null)
fun discard (constraint: Constraint, occ: ConstraintOccurrence) = PartialMatch(this, null, Pair(constraint, occ))
@ -85,48 +86,62 @@ class PartialMatch(val rule: Rule) {
fun isGuardSatisfied() : Boolean = true
fun matches(): Boolean {
return Unification.unify(this.toMatchTerm(), this.rule.toMatchTerm()).isSuccessful
val subst = Unification.unify(PartialMatchTerm(this), RuleTerm(this.rule))
if (!subst.isSuccessful) return false
// only one parameter of the unification can contain variables,
// thus triangular form never has variables on the right hand side
this.logicalContext = object: LogicalContext {
val var2val = subst.bindings().map { b ->
(b.`var`().symbol() as LogicalPattern).to(b.term().toValue()) }.toMap()
override fun valueFor(logicalPattern: LogicalPattern): Any? = var2val[logicalPattern]
}
return true
}
fun logicalContext(): LogicalContext = logicalContext ?: throw IllegalStateException("no logical context")
}
/**
* True iff the constraint matches the occurrence.
*/
fun Constraint.matches(that: ConstraintOccurrence): Boolean {
return Unification.unify(this.toMatchTerm(), that.toMatchTerm()).isSuccessful
return Unification.unify(ConstraintTerm(this), ConstraintOccurrenceTerm(that)).isSuccessful
}
fun PartialMatch.toMatchTerm(): Term = PartialMatchTerm(this)
class PartialMatchTerm(pm : PartialMatch) :
Function(pm.rule.tag(), pm.occurrences().map { co -> co.toMatchTerm() }) {}
fun Rule.toMatchTerm(): Term = RuleTerm(this)
/** Function term with arguments == constraints converted to terms. May contain variables. */
class RuleTerm(rule: Rule) :
Function(rule.tag(), (rule.headKept() + rule.headReplaced()).map { c -> ConstraintTerm(c) }) {}
fun Constraint.toMatchTerm(): Term = ConstraintTerm(this)
/** Function term with arguments == constraint arguments converted to terms.
* LogicalPattern arguments are term variables.
* Everything else is either a term or a constant wrapping the value. */
class ConstraintTerm(constraint: Constraint) :
Function(constraint.symbol(),
constraint.arguments().map { a -> if (a is LogicalPattern) Variable(a) else Constant(a!!) }) {}
constraint.arguments().map { arg -> if (arg is LogicalPattern) Variable(arg) else asTerm(arg) }) {}
fun ConstraintOccurrence.toMatchTerm(): Term = ConstraintOccurrenceTerm(this)
/** Function term with arguments == terms corresponding to constraint occurrences. Never contains variables. */
class PartialMatchTerm(pm : PartialMatch) :
Function(pm.rule.tag(), pm.occurrences().map { co -> ConstraintOccurrenceTerm(co) }) {}
/** Function term with arguments == constraint occurrence arguments converted to terms.
* Logical arguments are either terms/values (bound), or constants wrapping the logical itself (unbound).
* Everything else is either a term or a constant wrapping the value.
* Never contains variable terms. */
class ConstraintOccurrenceTerm(occurrence: ConstraintOccurrence) :
Function(occurrence.constraint().symbol(),
occurrence.arguments().map { co -> asTerm(co) }) {}
occurrence.arguments().map { arg -> if (arg is Logical<*>) arg.toTerm() else asTerm(arg) }) {}
fun asTerm(arg: Any?): Term {
return when(arg) {
is Logical<*> -> arg.toTerm()
is Term -> arg
else -> Constant(arg!!)
}
}
fun asTerm(arg: Any?): Term = if (arg is Term) arg else Constant(arg!!)
fun Logical<*>.toTerm(): Term = if (isBound) asTerm(findRoot().value()) else Constant(findRoot().value())
fun Logical<*>.toTerm(): Term = if (isBound) asTerm(findRoot().value()) else Constant(findRoot())
fun Term.toValue(): Any? = if (this is Constant) this.symbol() else this
abstract class TermImpl(val symbol: Any) : Term {

View File

@ -1,8 +1,9 @@
import jetbrains.mps.logic.reactor.constraint.*
import jetbrains.mps.logic.reactor.logical.NamingContext
import jetbrains.mps.logic.reactor.core.HandlingContext
import jetbrains.mps.logic.reactor.logical.LogicalContext
import jetbrains.mps.logic.reactor.logical.LogicalPattern
import jetbrains.mps.logic.reactor.rule.Rule
import jetbrains.mps.logic.reactor.rule.RuleBuilder
import jetbrains.mps.unification.Term
import java.util.*
/**
@ -11,7 +12,6 @@ import java.util.*
class Program(val env: Environment, val rules: List<Rule>) {
fun occurrenceFactory() : (Constraint) -> ConstraintOccurrence = { cst -> TestOccurrence(cst) }
}
class Environment() {
@ -66,7 +66,7 @@ fun equals(left: Any, right: Any): ConjBuilder.() -> Unit = {
add(TestEqPredicate(left, right))
}
fun occurrence(id: String, vararg args: Any) : ConstraintOccurrence = TestOccurrence(id, * args)
fun occurrence(id: String, vararg args: Any) : ConstraintOccurrence = TestConstraintOccurrence(id, * args)
class RB(tag: String, val env: Environment?) : RuleBuilder(tag) {
@ -117,21 +117,6 @@ private fun buildConjunction(type: Class<out AndItem>,
return conjBuilder
}
private data class TestOccurrence(val arguments : List<Any>, val constraint : Constraint) : ConstraintOccurrence {
constructor(id: String, vararg args: Any) :
this(listOf(* args), TestConstraint(ConstraintSymbol.symbol(id, args.size))) {}
constructor(constraint: Constraint) : this(constraint.arguments().map { it ?: TODO() }, constraint) {}
override fun constraint(): Constraint = constraint
override fun arguments(): Collection<Any> = arguments
override fun toString(): String = "#${constraint().symbol()}(${arguments().joinToString()})"
}
private data class TestConstraint(val symbol: ConstraintSymbol, val arguments: List<Any>) : Constraint {
constructor(symbol: ConstraintSymbol, vararg args: Any) : this(symbol, listOf(* args)) {}
@ -142,6 +127,31 @@ private data class TestConstraint(val symbol: ConstraintSymbol, val arguments: L
override fun argumentTypes(): List<Class<*>> = arguments.map { arg -> arg.javaClass }
override fun occurrence(context: LogicalContext): ConstraintOccurrence {
return TestConstraintOccurrence(this,
arguments.map { arg -> if (arg is LogicalPattern) context.valueFor(arg) else arg }.toList())
}
override fun toString(): String = "${symbol()}(${arguments().joinToString()})"
}
private data class TestConstraintOccurrence(val constraint: Constraint, val arguments: List<Any>, val id: Int) : ConstraintOccurrence {
companion object {
val random = Random()
}
constructor(constraint: Constraint, arguments: List<Any>) :
this(constraint, arguments, random.nextInt()) {}
constructor(id: String, vararg args: Any) :
this(TestConstraint(ConstraintSymbol.symbol(id, args.size)), listOf(* args), random.nextInt()) {}
override fun constraint(): Constraint = constraint
override fun arguments(): Collection<Any> = arguments
override fun toString(): String = "#${constraint().symbol()}(${arguments().joinToString()})#${id}"
}

View File

@ -21,7 +21,7 @@ class TestHandler {
init(PredicateSymbol("equals",2), JavaPredicateSymbol.EXPRESSION0, JavaPredicateSymbol.EXPRESSION1, JavaPredicateSymbol.EXPRESSION2, JavaPredicateSymbol.EXPRESSION3) }
fun Program.handler(vararg occurrences: ConstraintOccurrence): Handler =
Handler(sessionSolver(env.expressionSolver, env.equalsSolver), rules, occurrenceFactory(), listOf(* occurrences))
Handler(sessionSolver(env.expressionSolver, env.equalsSolver), rules, listOf(* occurrences))
companion object {
@BeforeClass @JvmStatic fun setup() {
@ -45,7 +45,9 @@ class TestHandler {
))
).run {
handler().apply { process(occurrence("main")) }.let { rh ->
assertEquals(setOf(occurrence("main"), occurrence("foo")), rh.occurrences())
assertEquals(
setOf(ConstraintSymbol("main", 0), ConstraintSymbol("foo", 0)),
rh.occurrences().map { it.constraint().symbol() }.toSet())
}
}
}
@ -72,7 +74,9 @@ class TestHandler {
))
).run {
handler().apply { process(occurrence("main")) }.let { rh ->
assertEquals(setOf(occurrence("bar"), occurrence("foo")), rh.occurrences())
assertEquals(
setOf(ConstraintSymbol("bar", 0), ConstraintSymbol("foo", 0)),
rh.occurrences().map { it.constraint().symbol() }.toSet())
}
}
}
@ -153,7 +157,9 @@ class TestHandler {
))
).run {
handler().apply { process(occurrence("main")) }.let { rh ->
assertEquals(setOf(occurrence("main"), occurrence("next")), rh.occurrences())
assertEquals(
setOf(ConstraintSymbol("main", 0), ConstraintSymbol("next", 0)),
rh.occurrences().map { it.constraint().symbol() }.toSet())
assertEquals("expected", test)
}
}

View File

@ -16,9 +16,9 @@ class TestMatcher {
fun Program.matcher(vararg occurrence: ConstraintOccurrence): Matcher {
val stored = occurrence.toList()
return object : Matcher(rules) {
override fun findOccurrences(constraint: Constraint, predicate: (ConstraintOccurrence) -> Boolean):
override fun findOccurrences(constraint: Constraint, acceptable: (ConstraintOccurrence) -> Boolean):
Iterable<ConstraintOccurrence> =
stored.filter { co -> constraint.matches(co) && predicate(co) }
stored.filter { co -> constraint.matches(co) && acceptable(co) }
}
}
@ -209,9 +209,18 @@ class TestMatcher {
assertFalse(matches.any())
}
// same parameter -- 4 matches (all permutations)
matcher(occurrence("foo", 42)).lookupMatches(occurrence("foo", 42)).let { matches ->
assertEquals(4, matches.count())
assertTrue(matches.all{ m -> m.occurrences().toSet().size == 2 })
}
matcher(occurrence("foo", 42)).lookupMatches(occurrence("foo", 16)).let { matches ->
assertEquals(2, matches.count())
assertEquals(listOf("main1", "main1"), matches.map { m -> m.rule.tag() })
matches.map { m -> setOf(M, N).map { lp -> m.logicalContext().valueFor(lp) } }.forEach { vals ->
assertEquals<Any>(setOf(42, 16), vals.toSet())
}
assertTrue(matches.all{ m -> m.occurrences().toSet().size == 2 })
}
@ -219,15 +228,19 @@ class TestMatcher {
x.find().value = 123
y.find().value = 456
matcher().lookupMatches(occurrence("foo", x)).let { matches ->
assertFalse(matches.any())
}
matcher(occurrence("foo", x)).lookupMatches(occurrence("foo", y)).let { matches ->
assertEquals(2, matches.count())
assertEquals(listOf("main1", "main1"), matches.map { m -> m.rule.tag() })
assertTrue(matches.all{ m -> m.occurrences().toSet().size == 2 })
}
val (v, w) = logical("v", "w")
w.find().union(v)
matcher(occurrence("foo", v)).lookupMatches(occurrence("foo", w)).let { matches ->
assertEquals(4, matches.count())
assertTrue(matches.all{ m -> m.occurrences().toSet().size == 2 })
}
}
}
}