diff --git a/reactor/API/src/jetbrains/mps/logic/reactor/program/Handler.java b/reactor/API/src/jetbrains/mps/logic/reactor/program/Handler.java index 8942af80..9bad4740 100644 --- a/reactor/API/src/jetbrains/mps/logic/reactor/program/Handler.java +++ b/reactor/API/src/jetbrains/mps/logic/reactor/program/Handler.java @@ -7,7 +7,7 @@ public abstract class Handler { public abstract String name(); - public abstract ConstraintSymbol primarySymbol(); + public abstract Iterable primarySymbols(); public abstract Iterable rules(); 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 b3b9e0a8..ca5d1dbb 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleIndex.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleIndex.kt @@ -33,35 +33,30 @@ class RuleIndex : Iterable { fun byTag(tag: String): Rule? = tag2rule[tag] - fun forOccurrence(occ: ConstraintOccurrence): Iterable { - val primary = primarySymbol2valueIndex.get(occ.constraint().symbol())?.select(occ) ?: emptyList() - val all = allSymbol2valueIndex.get(occ.constraint().symbol())?.select(occ) ?: emptyList() - return primary + all - } + fun forOccurrence(occ: ConstraintOccurrence): Iterable = + primarySymbol2valueIndex.get(occ.constraint().symbol())?.select(occ) ?: + (allSymbol2valueIndex.get(occ.constraint().symbol())?.select(occ) ?: emptyList()) override fun iterator(): Iterator = tag2rule.values.iterator() private fun buildIndex(handlers: Iterable) { + // first, init the primary symbols value index + handlers + .flatMap { h -> h.primarySymbols() } + .forEach { symbol -> primarySymbol2valueIndex.getOrPut(symbol) { ValueIndex(symbol) } } + for (h in handlers) { - val primaryValueIdx = h.primarySymbol()?.let { symbol -> - allSymbol2valueIndex.getOrPut(symbol) { ValueIndex(symbol) } - } + val hPrimSyms = h.primarySymbols().toSet() for (r in h.rules()) { - for (c in r.headKept()) { - if (c.symbol() == h.primarySymbol()) { - primaryValueIdx?.update(r, c) - - } else if (h.primarySymbol() == null) { - allSymbol2valueIndex.getOrPut(c.symbol()) { ValueIndex(c.symbol()) }.update(r, c) + for (c in (r.headKept() + r.headReplaced())) { + val symbol = c.symbol() + if (symbol in hPrimSyms) { + primarySymbol2valueIndex[symbol]?.update(r, c) } - } - for (c in r.headReplaced()) { - if (c.symbol() == h.primarySymbol()) { - primaryValueIdx?.update(r, c) - - } else if (h.primarySymbol() == null) { - allSymbol2valueIndex.getOrPut(c.symbol()) { ValueIndex(c.symbol()) }.update(r, c) + else if (hPrimSyms.isEmpty()) { + allSymbol2valueIndex.getOrPut(symbol) { ValueIndex(symbol) }.update(r, c) } + // else ignore the constraint -- it's not meant to be processed by this handler, as it is not a primary } } } diff --git a/reactor/Test/src/program/MockProgram.kt b/reactor/Test/src/program/MockProgram.kt index f0c4cd3d..11f8c426 100644 --- a/reactor/Test/src/program/MockProgram.kt +++ b/reactor/Test/src/program/MockProgram.kt @@ -27,7 +27,7 @@ class ProgramBuilder(val registry: ConstraintRegistry) { } -open class HandlerBuilder(val name: String, val primary: ConstraintSymbol?) { +open class HandlerBuilder(val name: String, val primary: Iterable) { val rules = ArrayList() fun appendRule(rule: Rule) { @@ -61,12 +61,12 @@ open class RuleBuilder(val tag: String) { class MockHandler( val name: String, - val primary: ConstraintSymbol?, + val primary: Iterable, val rules: List) : Handler() { override fun name(): String = name - override fun primarySymbol(): ConstraintSymbol? = primary + override fun primarySymbols(): Iterable = primary override fun rules(): Iterable = rules } diff --git a/reactor/Test/test/RulesHelper.kt b/reactor/Test/test/RulesHelper.kt index a0cf1e37..9414e75c 100644 --- a/reactor/Test/test/RulesHelper.kt +++ b/reactor/Test/test/RulesHelper.kt @@ -34,7 +34,7 @@ fun programWithRules(pb: ProgramBuilder, vararg ruleBuilders : Environment.() -> } private fun programWithRules(env: Environment, ruleBuilders: Array Rule>): Builder { - return builder(env, arrayOf(handler("test", null, * ruleBuilders))) + return builder(env, arrayOf(handler("test", emptyList(), * ruleBuilders))) } fun programWithHandlers(vararg handlerBuilders : Environment.() -> Handler): Builder { @@ -51,7 +51,7 @@ private fun builder(env: Environment, handlerBlocks: Array return Builder(env, handlers) } -fun handler(name: String, primary: ConstraintSymbol?, vararg ruleBlocks: Environment.() -> Rule): Environment.() -> Handler = { +fun handler(name: String, primary: Iterable, vararg ruleBlocks: Environment.() -> Rule): Environment.() -> Handler = { val hb = HandlerBuilder(name, primary) for (block in ruleBlocks) { hb.appendRule(this.block()) diff --git a/reactor/Test/test/TestMatcher.kt b/reactor/Test/test/TestMatcher.kt index ecc977e7..baac1496 100644 --- a/reactor/Test/test/TestMatcher.kt +++ b/reactor/Test/test/TestMatcher.kt @@ -184,13 +184,13 @@ class TestMatcher { @Test fun multipleHandlers() { programWithHandlers( - handler("handler1", ConstraintSymbol("foo", 0), + handler("handler1", listOf(ConstraintSymbol("foo", 0)), rule("main1", headKept( constraint("foo") )) ), - handler("handler2", ConstraintSymbol("bar", 0), + handler("handler2", listOf(ConstraintSymbol("bar", 0)), rule("main2", headKept( constraint("bar") diff --git a/reactor/Test/test/TestProgramBuilder.kt b/reactor/Test/test/TestProgramBuilder.kt index 2f1eb257..fb8fb878 100644 --- a/reactor/Test/test/TestProgramBuilder.kt +++ b/reactor/Test/test/TestProgramBuilder.kt @@ -39,7 +39,7 @@ class TestProgramBuilder { constraint("bar") ))).run { - programBuilder.addHandler(MockHandler("test", null, rules)) + programBuilder.addHandler(MockHandler("test", emptyList(), rules)) assertEquals(programBuilder.program("test").rules().count(), 1) assertEquals(programBuilder.program("test").rules().count(), 1) } @@ -66,7 +66,7 @@ class TestProgramBuilder { constraint("blah") ))).run { - programBuilder.addHandler(MockHandler("test", null, rules)) + programBuilder.addHandler(MockHandler("test", emptyList(), rules)) assertEquals(programBuilder.program("test").rules().count(), 2) } } @@ -82,7 +82,7 @@ class TestProgramBuilder { constraint("bar", "1") ))).run { - programBuilder.addHandler(MockHandler("test", null, rules)) + programBuilder.addHandler(MockHandler("test", emptyList(), rules)) } } }