From d15f6908d59b4487fed4a84b6dd5a268cc87b027 Mon Sep 17 00:00:00 2001 From: Fedor Isakov Date: Thu, 10 Dec 2015 12:56:45 +0100 Subject: [PATCH] Initial support for logicals. Small refactoring in tests. --- .../mps/logic/reactor/core/RuleHandler.kt | 9 +- .../reactor/predicate/ReactorSessionSolver.kt | 17 +- reactor/Test/test/JavaExpressionHelper.kt | 104 +++++++++ reactor/Test/test/LogicalHelper.kt | 161 +++++++++++++ reactor/Test/test/Rules.kt | 219 ------------------ reactor/Test/test/RulesHelper.kt | 141 +++++++++++ reactor/Test/test/TestBasicProgram.kt | 10 +- reactor/Test/test/TestPlanningSession.kt | 10 +- reactor/Test/test/TestRuleHandler.kt | 82 ++++++- 9 files changed, 500 insertions(+), 253 deletions(-) create mode 100644 reactor/Test/test/JavaExpressionHelper.kt create mode 100644 reactor/Test/test/LogicalHelper.kt delete mode 100644 reactor/Test/test/Rules.kt create mode 100644 reactor/Test/test/RulesHelper.kt diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleHandler.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleHandler.kt index 85577762..f3a9e2f7 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleHandler.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleHandler.kt @@ -65,14 +65,7 @@ class RuleHandler { private fun activate(constraint: Constraint): ConstraintOccurrence = occurrenceFactory(constraint) private fun tellPredicate(predicate: Predicate) { - when (predicate.symbol()) { - is JavaPredicateSymbol -> evalJava(predicate) - else -> TODO() - } - } - - private fun evalJava(expr: Predicate) { - sessionSolver.tell(expr.symbol(), * expr.arguments().toTypedArray()) + sessionSolver.tell(predicate.symbol(), * predicate.arguments().toTypedArray()) } fun lookupMatches(occ: ConstraintOccurrence): Iterable { diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/predicate/ReactorSessionSolver.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/predicate/ReactorSessionSolver.kt index 52e3d758..84c42346 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/predicate/ReactorSessionSolver.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/predicate/ReactorSessionSolver.kt @@ -6,17 +6,7 @@ import jetbrains.mps.logic.reactor.constraint.* * @author Fedor Isakov */ -class ReactorSessionSolver(val expressionSolver: Queryable) : SessionSolver() { - - constructor() : this(object: Queryable { - override fun ask(predicateSymbol: PredicateSymbol?, vararg args: Any?): Boolean { - throw UnsupportedOperationException() - } - - override fun tell(symbol: Symbol?, vararg args: Any?) { - throw UnsupportedOperationException() - } - }) +class ReactorSessionSolver(val expressionSolver: Queryable, val equalsSolver: Queryable) : SessionSolver() { override fun solverClass(predicateSymbol: PredicateSymbol?): Class? { throw UnsupportedOperationException() @@ -24,8 +14,9 @@ class ReactorSessionSolver(val expressionSolver: Queryable) : SessionSolver() { override fun registerSymbol(predicateSymbol: PredicateSymbol, computingTracer: ComputingTracer?) { when (predicateSymbol) { - is JavaPredicateSymbol -> registerSolver(predicateSymbol, expressionSolver) - else -> throw UnsupportedOperationException("not implemented") + is JavaPredicateSymbol -> registerSolver(predicateSymbol, expressionSolver) + PredicateSymbol("equals", 2) -> registerSolver(predicateSymbol, equalsSolver) + else -> throw UnsupportedOperationException("not implemented") } } diff --git a/reactor/Test/test/JavaExpressionHelper.kt b/reactor/Test/test/JavaExpressionHelper.kt new file mode 100644 index 00000000..b518dd5a --- /dev/null +++ b/reactor/Test/test/JavaExpressionHelper.kt @@ -0,0 +1,104 @@ +import jetbrains.mps.logic.reactor.constraint.* +import java.util.* + +/** + * @author Fedor Isakov + */ + + +fun expression(body: () -> Boolean): ConjBuilder.() -> Unit = { + add(TestJavaPredicate(JavaPredicateSymbol(1),body)) +} + +fun expression(body: (Any) -> Boolean, vararg args: Any): ConjBuilder.() -> Unit = { + add(TestJavaPredicate(JavaPredicateSymbol(2), body, * args)) +} + +fun expression(body: (Any, Any) -> Boolean, vararg args: Any): ConjBuilder.() -> Unit = { + add(TestJavaPredicate(JavaPredicateSymbol(3), body, * args)) +} + +fun expression(body: (Any, Any, Any) -> Boolean, vararg args: Any): ConjBuilder.() -> Unit = { + add(TestJavaPredicate(JavaPredicateSymbol(4), body, * args)) +} + +class ExpressionSolver : Queryable { + + override fun ask(predicateSymbol: PredicateSymbol, vararg args: Any): Boolean { + return javaPredicates[predicateSymbol]?.expr?.invoke(listOf(* args)) ?: + ERROR("no such symbol $predicateSymbol") + } + + override fun tell(symbol: Symbol, vararg args: Any) { + when (symbol) { + is JavaPredicateSymbol -> javaPredicates[args[0]]?.expr?.invoke(listOf(* args).drop(1)) + else -> ERROR("uknown symbol $symbol") + } + } + + fun addMaybeJavaPredicate(item: AndItem) { + if (item is TestJavaPredicate) { + addJavaPredicate(item) + } + } + + private val javaPredicates = HashMap() + + private fun addJavaPredicate(javaPredicate: TestJavaPredicate) { + javaPredicates[javaPredicate.args[0]] = javaPredicate + } + + private fun ERROR(msg: String) : Nothing = throw IllegalArgumentException(msg) +} + +private interface JavaExpression { + fun invoke(args: List): Boolean +} + +private class JavaExpression0(val code: () -> Boolean) : JavaExpression { + override fun invoke(args: List): Boolean { + if (args.size != 0) throw IllegalArgumentException("arity mismatch") + return code() + } +} + +private class JavaExpression1(val code: (Any) -> Boolean) : JavaExpression { + override fun invoke(args: List): Boolean { + if (args.size != 1) throw IllegalArgumentException("arity mismatch") + return code(args[0]) + } +} + +private class JavaExpression2(val code: (Any, Any) -> Boolean) : JavaExpression { + override fun invoke(args: List): Boolean { + if (args.size != 2) throw IllegalArgumentException("arity mismatch") + return code(args[0], args[1]) + } +} + +private class JavaExpression3(val code: (Any, Any, Any) -> Boolean) : JavaExpression { + override fun invoke(args: List): Boolean { + if (args.size != 3) throw IllegalArgumentException("arity mismatch") + return code(args[0], args[1], args[2]) + } +} + +private data class TestJavaPredicate(val symbol: JavaPredicateSymbol, val expr: JavaExpression, val args: List) : Predicate { + + constructor(symbol: JavaPredicateSymbol, code: () -> Boolean, vararg args: Any) : + this(symbol, JavaExpression0(code), listOf(code.toString()) + listOf(* args)) {} + + constructor(symbol: JavaPredicateSymbol, code: (Any) -> Boolean, vararg args: Any) : + this(symbol, JavaExpression1(code), listOf(code.toString()) + listOf(* args)) {} + + constructor(symbol: JavaPredicateSymbol, code: (Any, Any) -> Boolean, vararg args: Any) : + this(symbol, JavaExpression2(code), listOf(code.toString()) + listOf(* args)) {} + + constructor(symbol: JavaPredicateSymbol, code: (Any, Any, Any) -> Boolean, vararg args: Any) : + this(symbol, JavaExpression3(code), listOf(code.toString()) + listOf(* args)) {} + + override fun arguments(): List = args + + override fun symbol(): PredicateSymbol = symbol + +} diff --git a/reactor/Test/test/LogicalHelper.kt b/reactor/Test/test/LogicalHelper.kt new file mode 100644 index 00000000..fc156601 --- /dev/null +++ b/reactor/Test/test/LogicalHelper.kt @@ -0,0 +1,161 @@ +import jetbrains.mps.logic.reactor.constraint.Predicate +import jetbrains.mps.logic.reactor.constraint.PredicateSymbol +import jetbrains.mps.logic.reactor.constraint.Queryable +import jetbrains.mps.logic.reactor.constraint.Symbol +import jetbrains.mps.logic.reactor.logical.ILogical +import jetbrains.mps.logic.reactor.logical.NamingContext + +/** + * @author Fedor Isakov + */ + + +fun logical(name: String) = TestLogical(name) +fun logical(name1: String, name2: String) = Pair(TestLogical(name1), TestLogical(name2)) +fun logical(name1: String, name2: String, name3: String) = Triple(TestLogical(name1), TestLogical(name2), TestLogical(name3)) +fun setValue(logical: Any, value: Any) { (logical as TestLogical).find().value = value } +fun getValue(logical: Any) = (logical as TestLogical).find().value() + +class EqualsSolver : Queryable { + override fun ask(predicateSymbol: PredicateSymbol?, vararg args: Any?): Boolean { + if (args.size != 2) ERROR("arity mismatch") + val left = args[0] + val right = args[1] + + return if (left is TestLogical<*> && right is TestLogical<*>) { + ask_logical_logical(left, right) + } + else if (left is TestLogical<*>) { + ask_logical_value(left, right) + } + else if (right is TestLogical<*>) { + ask_value_logical(left, right) + } + else { + ask_value_value(left, right) + } + } + + override fun tell(symbol: Symbol, vararg args: Any) { + if (args.size != 2) ERROR("arity mismatch") + val left = args[0] + val right = args[1] + if (left is TestLogical<*> && right is TestLogical<*>) { + tell_logical_logical(left, right) + } + else if (left is TestLogical<*>) { + tell_logical_value(left, right) + } + else if (right is TestLogical<*>) { + tell_value_logical(left, right) + } + else { + tell_value_value(left, right) + } + } + + fun ask_logical_logical(left: TestLogical<*>, right: TestLogical<*>): Boolean { + return left.isBound && right.isBound && left.findRoot().value() == right.findRoot().value() + } + + fun ask_logical_value(left: TestLogical<*>, right: Any?): Boolean { + return left.isBound && left.findRoot().value() == right + } + + fun ask_value_logical(left: Any?, right: TestLogical<*>): Boolean { + return right.isBound && right.findRoot().value() == left + } + + fun ask_value_value(left: Any?, right: Any?): Boolean { + return left == right + } + + fun tell_logical_logical(left: TestLogical<*>, right: TestLogical<*>) { + if (left.isBound && right.isBound) { + check (left.find().value == right.find().value) + left.union(right) + } + else if (left.isBound) { + right.union(left) + } + else if (right.isBound) { + left.union(right) + } + else { + left.union(right) + } + } + + fun tell_logical_value(left: TestLogical<*>, right: Any?) { + if (left.isBound) { + check(left.find().value == right) + } + else { + // TODO hack! + (left.find() as TestLogical).value = right + } + } + + fun tell_value_logical(left: Any?, right: TestLogical<*>) { + if (right.isBound) { + check(right.find().value == left) + } + else { + // TODO: hack! + (right.find() as TestLogical).value = left + } + } + fun tell_value_value(left: Any?, right: Any?) { + check(left == right) + } + + private fun check(condition: Boolean) { + if (!condition) throw IllegalStateException() + } + + private fun ERROR(msg: String) : Nothing = throw IllegalArgumentException(msg) +} + + +data class TestLogical(val name: String, var value: T?, var parent: TestLogical?) : ILogical { + + constructor(name: String) : this(name, null, null) {} + + override fun name(): String = name + + override fun name(namingContext: NamingContext?): String? { + throw UnsupportedOperationException() + } + + override fun findRoot(): ILogical = find() + + override fun value(): T? = value + + override fun isBound(): Boolean = find().value != null + + override fun isWildcard(): Boolean { + throw UnsupportedOperationException() + } + + fun find(): TestLogical { + val tmp = parent + if (tmp == null) return this + else { + val root = tmp.find() + this.parent = root + return root + } + } + + fun union(other: TestLogical) { + if (find() != other.find()) find().parent = other + } +} + +data class TestEqPredicate(val left: Any, val right: Any) : Predicate { + + override fun arguments(): List = listOf(left, right) + + override fun symbol(): PredicateSymbol = PredicateSymbol("equals", 2) + +} \ No newline at end of file diff --git a/reactor/Test/test/Rules.kt b/reactor/Test/test/Rules.kt deleted file mode 100644 index 5f6851ab..00000000 --- a/reactor/Test/test/Rules.kt +++ /dev/null @@ -1,219 +0,0 @@ -import jetbrains.mps.logic.reactor.constraint.* -import jetbrains.mps.logic.reactor.rule.Rule -import jetbrains.mps.logic.reactor.rule.RuleBuilder -import java.util.* - -/** - * @author Fedor Isakov - */ - -fun program(vararg ruleBuilders : Environment.() -> Rule): Program { - val env = Environment() - val rules = ArrayList() - with (env) { - for (rb in ruleBuilders) { - rules.add(rb()) - } - } - return Program(env, rules) -} - -fun rule(tag: String, vararg component:RB.() -> Unit): Environment.() -> Rule = { - rule(tag, this, * component) -} - -fun rule(tag: String, env: Environment, vararg component:RB.() -> Unit): Rule { - val rb = RB(tag, env) - for (cmp in component) { - rb.cmp() - } - return rb.toRule() -} - -fun headKept(vararg content : ConjBuilder.() -> Unit): RB.() -> Unit = { - appendHeadKept( * buildConjunction(Constraint::class.java, env, content).toArray()) -} -fun headReplaced(vararg content : ConjBuilder.() -> Unit): RB.() -> Unit = { - appendHeadReplaced( * buildConjunction(Constraint::class.java, env, content).toArray()) -} -fun guard(vararg content : ConjBuilder.() -> Unit): RB.() -> Unit = { - appendGuard( * buildConjunction(Predicate::class.java, env, content).toArray()) -} -fun body(vararg content : ConjBuilder.() -> Unit): RB.() -> Unit = { - appendBody( * buildConjunction(AndItem::class.java, env, content).toArray()) -} - -fun constraint(id: String, vararg args: Any): ConjBuilder.() -> Unit = { - add(TestConstraint(ConstraintSymbol.symbol(id, args.size), * args)) -} - -fun expression(body: () -> Boolean): ConjBuilder.() -> Unit = { - add(TestJavaPredicate(JavaPredicateSymbol(1),body)) -} - -fun expression(body: (Any) -> Boolean, vararg args: Any): ConjBuilder.() -> Unit = { - add(TestJavaPredicate(JavaPredicateSymbol(2), body, * args)) -} - -fun expression(body: (Any, Any) -> Boolean, vararg args: Any): ConjBuilder.() -> Unit = { - add(TestJavaPredicate(JavaPredicateSymbol(3), body, * args)) -} - -fun expression(body: (Any, Any, Any) -> Boolean, vararg args: Any): ConjBuilder.() -> Unit = { - add(TestJavaPredicate(JavaPredicateSymbol(4), body, * args)) -} - -class Program(val env: Environment, val rules: List) { - fun occurrenceFactory() : (Constraint) -> ConstraintOccurrence = { cst -> TestOccurrence(cst) } -} - -class Environment() { - - val javaPredicates = HashMap() - - fun expressionSolver(): Queryable = object : Queryable { - - override fun ask(predicateSymbol: PredicateSymbol, vararg args: Any): Boolean { - return javaPredicates[predicateSymbol]?.expr?.invoke(listOf(* args)) ?: - ERROR("no such symbol $predicateSymbol") - } - - override fun tell(symbol: Symbol, vararg args: Any) { - when (symbol) { - is JavaPredicateSymbol -> javaPredicates[args[0]]?.expr?.invoke(listOf(* args).drop(1)) - else -> ERROR("uknown symbol $symbol") - } - } - - private fun ERROR(msg: String) : Nothing = throw IllegalArgumentException(msg) - } - - fun addJavaPredicate(javaPredicate: TestJavaPredicate) { - javaPredicates[javaPredicate.args[0]] = javaPredicate - } -} - -class RB(tag: String, val env: Environment?) : RuleBuilder(tag) {} - -class ConjBuilder { - val constraints = ArrayList() - val type: Class - val _env: Environment? - val env: Environment - get() { return _env ?: throw IllegalStateException("no enviroment") } - - constructor(type: Class, env: Environment?) { - this.type = type - this._env = env - } - - fun add(item: AndItem): Unit { - if (!type.isAssignableFrom(item.javaClass)) - throw IllegalArgumentException("unexpected constraint class '${item.javaClass}'") - constraints.add(item) - if (item is TestJavaPredicate) { - env.addJavaPredicate(item) - } - } - fun toArray(): Array = - if (Constraint::class.java.isAssignableFrom(type)) - Array(constraints.size) { - constraints.get(it) as Constraint - } as Array - else - Array(constraints.size) { - constraints.get(it) - } as Array -} - -private fun buildConjunction(type: Class, - env: Environment?, - content: Array Unit>): ConjBuilder -{ - var conjBuilder = ConjBuilder(type, env) - for (c in content) { - conjBuilder.c() - } - return conjBuilder -} - -data class TestOccurrence(val arguments : List, 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(), constraint) {} - - override fun constraint(): Constraint = constraint - - override fun arguments(): List = arguments - - override fun toString(): String = "#${constraint().symbol()}(${arguments().joinToString()})" - -} - -data class TestConstraint(val symbol: ConstraintSymbol, val arguments: List) : Constraint { - - constructor(symbol: ConstraintSymbol, vararg args: Any) : this(symbol, listOf(* args)) {} - - override fun arguments(): List = arguments - - override fun symbol(): ConstraintSymbol = symbol - - override fun argumentTypes(): List> = arguments.map { arg -> arg.javaClass } - - override fun toString(): String = "${symbol()}(${arguments().joinToString()})" - -} - -interface JavaExpression { - fun invoke(args: List): Boolean -} - -class JavaExpression0(val code: () -> Boolean) : JavaExpression { - override fun invoke(args: List): Boolean { - if (args.size != 0) throw IllegalArgumentException("arity mismatch") - return code() - } -} - -class JavaExpression1(val code: (Any) -> Boolean) : JavaExpression { - override fun invoke(args: List): Boolean { - if (args.size != 1) throw IllegalArgumentException("arity mismatch") - return code(args[0]) - } -} - -class JavaExpression2(val code: (Any, Any) -> Boolean) : JavaExpression { - override fun invoke(args: List): Boolean { - if (args.size != 2) throw IllegalArgumentException("arity mismatch") - return code(args[0], args[1]) - } -} - -class JavaExpression3(val code: (Any, Any, Any) -> Boolean) : JavaExpression { - override fun invoke(args: List): Boolean { - if (args.size != 3) throw IllegalArgumentException("arity mismatch") - return code(args[0], args[1], args[2]) - } -} - -data class TestJavaPredicate(val symbol: JavaPredicateSymbol, val expr: JavaExpression, val args: List) : Predicate { - - constructor(symbol: JavaPredicateSymbol, code: () -> Boolean, vararg args: Any) : - this(symbol, JavaExpression0(code), listOf(code.toString()) + listOf(* args)) {} - - constructor(symbol: JavaPredicateSymbol, code: (Any) -> Boolean, vararg args: Any) : - this(symbol, JavaExpression1(code), listOf(code.toString()) + listOf(* args)) {} - - constructor(symbol: JavaPredicateSymbol, code: (Any, Any) -> Boolean, vararg args: Any) : - this(symbol, JavaExpression2(code), listOf(code.toString()) + listOf(* args)) {} - - constructor(symbol: JavaPredicateSymbol, code: (Any, Any, Any) -> Boolean, vararg args: Any) : - this(symbol, JavaExpression3(code), listOf(code.toString()) + listOf(* args)) {} - - override fun arguments(): List = args - - override fun symbol(): PredicateSymbol = symbol - -} diff --git a/reactor/Test/test/RulesHelper.kt b/reactor/Test/test/RulesHelper.kt new file mode 100644 index 00000000..c6f3dddb --- /dev/null +++ b/reactor/Test/test/RulesHelper.kt @@ -0,0 +1,141 @@ +import jetbrains.mps.logic.reactor.constraint.* +import jetbrains.mps.logic.reactor.logical.ILogical +import jetbrains.mps.logic.reactor.logical.NamingContext +import jetbrains.mps.logic.reactor.rule.Rule +import jetbrains.mps.logic.reactor.rule.RuleBuilder +import java.util.* + +/** + * @author Fedor Isakov + */ + + +class Program(val env: Environment, val rules: List) { + fun occurrenceFactory() : (Constraint) -> ConstraintOccurrence = { cst -> TestOccurrence(cst) } +} + +class Environment() { + val equalsSolver = EqualsSolver() + val expressionSolver = ExpressionSolver() +} + +fun program(vararg ruleBuilders : Environment.() -> Rule): Program { + val env = Environment() + val rules = ArrayList() + with (env) { + for (rb in ruleBuilders) { + rules.add(rb()) + } + } + return Program(env, rules) +} + +fun rule(tag: String, vararg component:RB.() -> Unit): Environment.() -> Rule = { + rule(tag, this, * component) +} + +fun rule(tag: String, env: Environment, vararg component:RB.() -> Unit): Rule { + val rb = RB(tag, env) + for (cmp in component) { + rb.cmp() + } + return rb.toRule() +} + +fun headKept(vararg content : ConjBuilder.() -> Unit): RB.() -> Unit = { + appendHeadKept( * buildConjunction(Constraint::class.java, env, content).toArray()) +} + +fun headReplaced(vararg content : ConjBuilder.() -> Unit): RB.() -> Unit = { + appendHeadReplaced( * buildConjunction(Constraint::class.java, env, content).toArray()) +} + +fun guard(vararg content : ConjBuilder.() -> Unit): RB.() -> Unit = { + appendGuard( * buildConjunction(Predicate::class.java, env, content).toArray()) +} + +fun body(vararg content : ConjBuilder.() -> Unit): RB.() -> Unit = { + appendBody( * buildConjunction(AndItem::class.java, env, content).toArray()) +} + +fun constraint(id: String, vararg args: Any): ConjBuilder.() -> Unit = { + add(TestConstraint(ConstraintSymbol(id, args.size), * args)) +} + +fun equals(left: Any, right: Any): ConjBuilder.() -> Unit = { + add(TestEqPredicate(left, right)) +} + +fun occurrence(id: String, vararg args: Any) : ConstraintOccurrence = TestOccurrence(id, * args) + +class RB(tag: String, val env: Environment?) : RuleBuilder(tag) {} + +class ConjBuilder { + val constraints = ArrayList() + val type: Class + val _env: Environment? + val env: Environment + get() { return _env ?: throw IllegalStateException("no enviroment") } + + constructor(type: Class, env: Environment?) { + this.type = type + this._env = env + } + + fun add(item: AndItem): Unit { + if (!type.isAssignableFrom(item.javaClass)) + throw IllegalArgumentException("unexpected constraint class '${item.javaClass}'") + constraints.add(item) + env.expressionSolver.addMaybeJavaPredicate(item) + } + + fun toArray(): Array = + if (Constraint::class.java.isAssignableFrom(type)) + Array(constraints.size) { + constraints.get(it) as Constraint + } as Array + else + Array(constraints.size) { + constraints.get(it) + } as Array +} + +private fun buildConjunction(type: Class, + env: Environment?, + content: Array Unit>): ConjBuilder +{ + var conjBuilder = ConjBuilder(type, env) + for (c in content) { + conjBuilder.c() + } + return conjBuilder +} + +private data class TestOccurrence(val arguments : List, 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(), constraint) {} + + override fun constraint(): Constraint = constraint + + override fun arguments(): List = arguments + + override fun toString(): String = "#${constraint().symbol()}(${arguments().joinToString()})" + +} + +private data class TestConstraint(val symbol: ConstraintSymbol, val arguments: List) : Constraint { + + constructor(symbol: ConstraintSymbol, vararg args: Any) : this(symbol, listOf(* args)) {} + + override fun arguments(): List = arguments + + override fun symbol(): ConstraintSymbol = symbol + + override fun argumentTypes(): List> = arguments.map { arg -> arg.javaClass } + + override fun toString(): String = "${symbol()}(${arguments().joinToString()})" + +} diff --git a/reactor/Test/test/TestBasicProgram.kt b/reactor/Test/test/TestBasicProgram.kt index 2634fe1b..81784bd5 100644 --- a/reactor/Test/test/TestBasicProgram.kt +++ b/reactor/Test/test/TestBasicProgram.kt @@ -1,3 +1,6 @@ +import jetbrains.mps.logic.reactor.constraint.PredicateSymbol +import jetbrains.mps.logic.reactor.constraint.Queryable +import jetbrains.mps.logic.reactor.constraint.Symbol import jetbrains.mps.logic.reactor.core.ReactorEvaluationSession import jetbrains.mps.logic.reactor.core.ReactorPlanningSession import jetbrains.mps.logic.reactor.predicate.ReactorSessionSolver @@ -26,8 +29,13 @@ class TestBasicProgram { } } + val dummySolver = object : Queryable { + override fun ask(predicateSymbol: PredicateSymbol?, vararg args: Any?): Boolean = TODO() + override fun tell(symbol: Symbol?, vararg args: Any?) = TODO() + } + @Before fun beforeTest() { - planningSession = PlanningSession.newSession("test", ReactorSessionSolver()) + planningSession = PlanningSession.newSession("test", ReactorSessionSolver(dummySolver, dummySolver)) evalConfig = EvaluationSession.newSession(planningSession) } diff --git a/reactor/Test/test/TestPlanningSession.kt b/reactor/Test/test/TestPlanningSession.kt index 521b6624..70212a43 100644 --- a/reactor/Test/test/TestPlanningSession.kt +++ b/reactor/Test/test/TestPlanningSession.kt @@ -1,3 +1,6 @@ +import jetbrains.mps.logic.reactor.constraint.PredicateSymbol +import jetbrains.mps.logic.reactor.constraint.Queryable +import jetbrains.mps.logic.reactor.constraint.Symbol import jetbrains.mps.logic.reactor.program.PlanningSession import jetbrains.mps.logic.reactor.rule.InvalidConstraintException import jetbrains.mps.logic.reactor.rule.InvalidRuleException @@ -23,8 +26,13 @@ class TestPlanningSession { } } + val dummySolver = object : Queryable { + override fun ask(predicateSymbol: PredicateSymbol?, vararg args: Any?): Boolean = TODO() + override fun tell(symbol: Symbol?, vararg args: Any?) = TODO() + } + @Before fun beforeTest() { - session = PlanningSession.newSession("test", ReactorSessionSolver()) + session = PlanningSession.newSession("test", ReactorSessionSolver(dummySolver, dummySolver)) } lateinit var session: PlanningSession diff --git a/reactor/Test/test/TestRuleHandler.kt b/reactor/Test/test/TestRuleHandler.kt index 16c9cf5d..d9696a64 100644 --- a/reactor/Test/test/TestRuleHandler.kt +++ b/reactor/Test/test/TestRuleHandler.kt @@ -1,8 +1,6 @@ -import jetbrains.mps.logic.reactor.constraint.ConstraintOccurrence -import jetbrains.mps.logic.reactor.constraint.JavaPredicateSymbol -import jetbrains.mps.logic.reactor.constraint.Queryable -import jetbrains.mps.logic.reactor.constraint.SessionSolver +import jetbrains.mps.logic.reactor.constraint.* import jetbrains.mps.logic.reactor.core.RuleHandler +import jetbrains.mps.logic.reactor.logical.ILogical import jetbrains.mps.logic.reactor.predicate.ReactorSessionSolver import org.junit.Before import org.junit.BeforeClass @@ -18,13 +16,12 @@ import kotlin.test.assertTrue class TestRuleHandler { - fun occurrence(id: String, vararg args: Any) : TestOccurrence = TestOccurrence(id, * args) - - fun sessionSolver(exprSolver: Queryable) : SessionSolver = - ReactorSessionSolver(exprSolver).apply { init(JavaPredicateSymbol.EXPRESSION0, JavaPredicateSymbol.EXPRESSION1, JavaPredicateSymbol.EXPRESSION2, JavaPredicateSymbol.EXPRESSION3) } + fun sessionSolver(exprSolver: Queryable, equalsSolver: Queryable) : SessionSolver = + ReactorSessionSolver(exprSolver, equalsSolver).apply { + init(PredicateSymbol("equals",2), JavaPredicateSymbol.EXPRESSION0, JavaPredicateSymbol.EXPRESSION1, JavaPredicateSymbol.EXPRESSION2, JavaPredicateSymbol.EXPRESSION3) } fun Program.handler(vararg occurrences: ConstraintOccurrence): RuleHandler = - RuleHandler(sessionSolver(env.expressionSolver()), rules, occurrenceFactory(), listOf(* occurrences)) + RuleHandler(sessionSolver(env.expressionSolver, env.equalsSolver), rules, occurrenceFactory(), listOf(* occurrences)) companion object { @BeforeClass @JvmStatic fun setup() { @@ -135,7 +132,7 @@ class TestRuleHandler { constraint("bar") ))).run { - val matches = handler(TestOccurrence("aux")).lookupMatches(occurrence("main")) + val matches = handler(occurrence("aux")).lookupMatches(occurrence("main")) assertFalse { matches.any { m -> m.isPartial() } } assertEquals(rules, matches.map { m -> m.rule }) @@ -252,9 +249,72 @@ class TestRuleHandler { ))).run { val result = handler().process(occurrence("main")) - assertEquals("value", test) + } + } + @Test + fun paramExpression() { + var test : String = "not initialized" + program( + rule("main", + headKept( + constraint("main") + ), + body( + expression ({ v -> test = v as String; true }, "value") + ))).run { + + val result = handler().process(occurrence("main")) + assertEquals("value", test) + } + } + + @Test + fun basicLogical() { + var test : String? = "not initialized" + val x = logical("x") + x.value = "expected" + program( + rule("main", + headKept( + constraint("main") + ), + body( + expression ({ v -> test = (v as ILogical).value(); true }, x) + ))).run { + + val result = handler().process(occurrence("main")) + assertEquals("expected", test) + } + } + + @Test + fun logicalCopy() { + var test : String? = "not initialized" + val (x,y) = logical("x", "y") + x.value = "expected" + program( + rule("main", + headKept( + constraint("main") + ), + body( + equals(x, y), + constraint("next") + )), + rule("aux", + headKept( + constraint("next") + ), + body( + expression ({ v -> test = getValue(v) as String; true }, y) + ))).run { + + handler().apply { process(occurrence("main")) }.run { + assertEquals(setOf(occurrence("main"), occurrence("next")), occurrences()) + } + assertEquals("expected", test) } }