Adjust nullability contract in occurrence matcher.

Term's symbol can't be null.
Also guarantee that when processing a match no null
values occur.
This commit is contained in:
Fedor Isakov 2020-02-28 17:27:54 +01:00
parent fec446f006
commit 3193a9c078
5 changed files with 46 additions and 30 deletions

View File

@ -34,7 +34,7 @@ interface OccurrenceMatcher {
fun matches(cst: Constraint, occ: ConstraintOccurrence): Boolean
fun match(left: Any?, right: Any?): Boolean
fun match(left: Any, right: Any): Boolean
}

View File

@ -26,12 +26,24 @@ import jetbrains.mps.logic.reactor.logical.MetaLogical
import jetbrains.mps.logic.reactor.program.Constraint
import jetbrains.mps.logic.reactor.util.assoc
import jetbrains.mps.unification.Term
import java.lang.NullPointerException
internal class OccurrenceMatcherImpl(val contextSubst: Subst? = null) : OccurrenceMatcher {
private var matchSubst : Subst? = null
private companion object {
// to be used instead of null for resolved logical w/o value
val NO_VALUE = object {}
}
override fun subst(): Subst = matchSubst ?: (contextSubst ?: emptySubst())
private var _matchSubst : Subst? = null
private var matchSubst : Subst
get() = _matchSubst ?: if (contextSubst != null) contextSubst else emptySubst()
set(value) {
this._matchSubst = value
}
override fun subst(): Subst = matchSubst
/**
* Matches constraint and occurrence.
@ -43,6 +55,8 @@ internal class OccurrenceMatcherImpl(val contextSubst: Subst? = null) : Occurren
if (cst.symbol() != occ.constraint().symbol()) return false
return zipWhileTrue(cst.arguments(), occ.arguments()) { cstarg, occarg ->
if (cstarg == null) throw NullPointerException("constraint argument can't be null")
if (occarg == null) return false
ptnMatchAny(cstarg, occarg)
}
}
@ -50,7 +64,7 @@ internal class OccurrenceMatcherImpl(val contextSubst: Subst? = null) : Occurren
/**
* Returns true for matching left and right parameters, false otherwise.
*/
override fun match(left: Any?, right: Any?): Boolean {
override fun match(left: Any, right: Any): Boolean {
return matchAny(left, right)
}
@ -58,20 +72,16 @@ internal class OccurrenceMatcherImpl(val contextSubst: Subst? = null) : Occurren
* Matches target against pattern.
* Recursively iterates terms.
* Respects substitutions for MetaLogical instances.
* Returns either new substitution on successful match, or null.
* Returns true on successful match, otherwise false.
*/
private fun ptnMatchAny(ptn: Any?, trg: Any?): Boolean =
private fun ptnMatchAny(ptn: Any, trg: Any): Boolean =
when (ptn) {
is MetaLogical<*> -> {
// recursion with existing substitution or new substitution
if (matchSubst == null) {
this.matchSubst = if (contextSubst != null) contextSubst else emptySubst()
}
if (matchSubst!!.containsKey(ptn))
matchSubst!![ptn].let { matchAny(it, trg) }
if (matchSubst.containsKey(ptn))
matchSubst[ptn]?.let { matchAny(it, trg) } ?: false
else
matchSubst!!.assoc(ptn, trg!!).also { matchSubst = it }.run { true }
matchSubst.assoc(ptn, trg).also { this.matchSubst = it }.run { true }
}
is Term ->
when {
@ -80,15 +90,14 @@ internal class OccurrenceMatcherImpl(val contextSubst: Subst? = null) : Occurren
}
else ->
when {
trg is Logical<*> -> ptnMatchAny(ptn, resolve(trg))
else ->
// compare two arbitrary values
trg is Logical<*> -> ptnMatchAny(ptn, resolve(trg))
else -> // compare two arbitrary values
(ptn == trg)
}
}
private fun matchAny(left: Any?, right: Any?): Boolean =
private fun matchAny(left: Any, right: Any): Boolean =
when (left) {
is Logical<*> ->
// match logical or its value
@ -100,15 +109,13 @@ internal class OccurrenceMatcherImpl(val contextSubst: Subst? = null) : Occurren
}
else ->
when {
right is Logical<*> -> matchAny(left, resolve(right))
right is Logical<*> -> matchAny(left, resolve(right))
else -> // compare two arbitrary values
(left == right)
}
}
private fun ptnMatchTerm(ptn: Term, trg: Any?): Boolean {
if (trg == null) return false
private fun ptnMatchTerm(ptn: Term, trg: Any): Boolean {
val trgval = resolve(trg)
if (!(trgval is Term)) return false
@ -124,9 +131,7 @@ internal class OccurrenceMatcherImpl(val contextSubst: Subst? = null) : Occurren
}
}
private fun matchTerm(left: Term, right: Any?): Boolean {
if (right == null) return false
private fun matchTerm(left: Term, right: Any): Boolean {
val rval = resolve(right)
if (!(rval is Term)) return false
@ -140,12 +145,11 @@ internal class OccurrenceMatcherImpl(val contextSubst: Subst? = null) : Occurren
}
}
private fun matchLogical(left: Logical<*>, right: Any?): Boolean =
private fun matchLogical(left: Logical<*>, right: Any): Boolean =
when {
right is Logical<*> ->
when {
left.isBound -> matchAny(left.findRoot().value(),
right.findRoot().value())
left.isBound -> matchAny(left.findRoot().value(), resolve(right))
left.findRoot() === right.findRoot() -> true // reference equality
else -> false
}
@ -153,11 +157,11 @@ internal class OccurrenceMatcherImpl(val contextSubst: Subst? = null) : Occurren
else -> false
}
private fun resolve(obj: Any?): Any? =
private fun resolve(obj: Any): Any =
when (obj) {
is LogicalOwner -> if (obj.logical().isBound) resolve(obj.logical()) else obj
is Logical<*> -> resolve(obj.findRoot().value())
is Term -> if (obj.`is`(Term.Kind.REF)) resolve(obj.get()) else obj
is Logical<*> -> if (obj.isBound) resolve(obj.findRoot().value()) else NO_VALUE
is Term -> if (obj.`is`(Term.Kind.REF)) resolve(obj.get()) else obj
else -> obj
}

View File

@ -16,6 +16,8 @@
package jetbrains.mps.unification;
import org.jetbrains.annotations.NotNull;
import java.util.Collection;
/**
@ -32,9 +34,11 @@ public interface Term extends Comparable<Term> {
/**
* A symbol corresponding to the function/constant name or the variable name.
* Must not be null.
* TODO: what is the meaning of symbol in case of a reference term?
* NOTE: two variable terms are treated as the same term if their symbols match.
*/
@NotNull
Object symbol();
/**

View File

@ -19,6 +19,7 @@ package jetbrains.mps.unification.test;
import jetbrains.mps.logic.reactor.logical.Logical;
import jetbrains.mps.logic.reactor.logical.LogicalOwner;
import jetbrains.mps.unification.Term;
import org.jetbrains.annotations.NotNull;
import java.util.Arrays;
import java.util.Collection;
@ -61,6 +62,7 @@ public abstract class MockTerm implements Term {
Term lookupTerm();
}
@NotNull
@Override
public Object symbol() {
return null;
@ -90,6 +92,7 @@ public abstract class MockTerm implements Term {
this.myArgs = Arrays.asList(arguments);
}
@NotNull
@Override
public Object symbol() {
return mySymbol;
@ -140,6 +143,7 @@ public abstract class MockTerm implements Term {
myName = name;
}
@NotNull
@Override
public Object symbol() {
return myName;
@ -174,6 +178,7 @@ public abstract class MockTerm implements Term {
this.mySymbol = symbol;
}
@NotNull
@Override
public Object symbol() {
return mySymbol;
@ -212,6 +217,7 @@ public abstract class MockTerm implements Term {
return myLogical;
}
@NotNull
@Override
public Object symbol() {
return myLogical;
@ -266,6 +272,7 @@ public abstract class MockTerm implements Term {
this.termLookup = termLookup;
}
@NotNull
@Override
public Object symbol() {
return "^";

View File

@ -414,6 +414,7 @@ public class SolverTests {
this.wrapped = term;
}
@NotNull
@Override
public Object symbol() {
return wrapped;